diff --git a/miles/ray/rollout/inference_controller.py b/miles/ray/rollout/inference_controller.py index e5f4c7dae00..3053b6b2338 100644 --- a/miles/ray/rollout/inference_controller.py +++ b/miles/ray/rollout/inference_controller.py @@ -21,6 +21,7 @@ requires_lock, with_lock, ) +from miles.utils.ft_utils.health_checker import ActivenessTracker from miles.utils.misc import SimpleTicker from miles.utils.workers.worker_provider.base import BaseWorkerProvider, CellInfo, StopWatchFn from miles.utils.workers.worker_provider.ray import RayWorkerProvider @@ -43,6 +44,7 @@ def __init__(self, args): self.rollout_id = -1 self.eval_fleet: EvalFleet | None = None self._watcher_disposers: list[StopWatchFn] = [] + self._health_checker_activeness = ActivenessTracker(active=True) self._ticker: SimpleTicker | None = None @lock_exempt @@ -50,7 +52,11 @@ async def init(self) -> None: if self.args.debug_train_only: return - self.servers = await create_rollout_servers(self.args, context_lock=self.context_lock) + self.servers = await create_rollout_servers( + self.args, + context_lock=self.context_lock, + global_health_checker_activeness=self._health_checker_activeness.get, + ) if self.args.eval_num_gpus > 0: self.eval_fleet = EvalFleet(self.args, srv=self.servers["eval"]) @@ -90,6 +96,9 @@ async def dispose(self): await disposer() self._watcher_disposers = [] + for srv in self.servers.values(): + await srv.dispose() + # -------------------------- offload/onload ----------------------------- # TODO may parallelly execute offload/onload across services @@ -267,25 +276,24 @@ async def _reconcile(self, cell_id: str, observed: CellInfo | None) -> None: @requires_lock async def _health_monitoring_pause(self) -> None: - self._assert_rollout_fault_tolerance_is_unsupported() + self._health_checker_activeness.bump_active(False) + await asyncio.gather( + *[ + cell.cancel_inflight_health_probe() + for srv in self.servers.values() + for cell in srv.server_cells.values() + ] + ) @requires_lock async def _health_monitoring_resume(self) -> None: - self._assert_rollout_fault_tolerance_is_unsupported() + self._health_checker_activeness.bump_active(True) @property @requires_lock def _rollout_ft_enabled(self) -> bool: return self.args.use_fault_tolerance and "rollout" in self.args.ft_components - @requires_lock - def _assert_rollout_fault_tolerance_is_unsupported(self) -> None: - if not self.args.debug_train_only and self._rollout_ft_enabled: - raise NotImplementedError( - "rollout fault tolerance is being rebuilt; health monitoring must pause before " - "get_updatable_engines snapshots the engines" - ) - @property @requires_lock def _server(self) -> RolloutServer | None: diff --git a/miles/ray/rollout/rollout_server.py b/miles/ray/rollout/rollout_server.py index 4c1957214b4..1c9917a894b 100644 --- a/miles/ray/rollout/rollout_server.py +++ b/miles/ray/rollout/rollout_server.py @@ -1,6 +1,7 @@ import asyncio import dataclasses import logging +from collections.abc import Callable from typing import Any from miles.backends.sglang_utils.sglang_api_client import SGLangApiClient @@ -9,6 +10,7 @@ from miles.ray.rollout.router_manager import wait_router_ready from miles.ray.rollout.server_cell import ServerCell, ServerCellMetadata from miles.utils.context_lock import ContextLock, enforce_lock_discipline, lock_exempt, requires_lock +from miles.utils.ft_utils.health_checker import ActiveAndEpoch from miles.utils.retry_utils import retry_until_deadline logger = logging.getLogger(__name__) @@ -17,7 +19,9 @@ WAIT_CELLS_MAX_DELAY_SECONDS = 5.0 -async def create_rollout_servers(args, context_lock: ContextLock) -> dict[str, "RolloutServer"]: +async def create_rollout_servers( + args, context_lock: ContextLock, global_health_checker_activeness: Callable[[], ActiveAndEpoch] +) -> dict[str, "RolloutServer"]: """Create rollout servers: one per model, each with its own router.""" assert args.sglang_router_ip is None, ( "external router mode was removed: miles always starts its own routers " @@ -43,6 +47,7 @@ async def create_rollout_servers(args, context_lock: ContextLock) -> dict[str, " router_port=router_addr.port, model_name=model_cfg.name, update_weights=model_cfg.update_weights, + global_health_checker_activeness=global_health_checker_activeness, expected_num_cells=model_cfg.num_server_cells, ) @@ -66,6 +71,9 @@ class RolloutServer: router_port: int | None = None model_name: str = "default" update_weights: bool = True + global_health_checker_activeness: Callable[[], ActiveAndEpoch] = lock_exempt( + lambda: ActiveAndEpoch(active=True, epoch=0) + ) expected_num_cells: int = 0 @property @@ -102,10 +110,15 @@ async def probe_and_mark_dead(self): async def add_cell(self, cell_meta: ServerCellMetadata): cell_id = cell_meta.cell_id assert cell_id not in self.server_cells - cell = ServerCell(args=self.args, router_api_client=self._router_api_client, meta=cell_meta) + cell = ServerCell( + args=self.args, + router_api_client=self._router_api_client, + meta=cell_meta, + global_health_checker_activeness=self.global_health_checker_activeness, + ) + self.server_cells[cell_id] = cell if not (self.args.colocate and cell_meta.needs_offload): await cell.init() - self.server_cells[cell_id] = cell @requires_lock async def remove_cell(self, cell_id: str): @@ -113,6 +126,11 @@ async def remove_cell(self, cell_id: str): await self.server_cells[cell_id].dispose() del self.server_cells[cell_id] + @requires_lock + async def dispose(self) -> None: + for cell_id in list(self.server_cells.keys()): + await self.remove_cell(cell_id) + @requires_lock async def offload(self, tags: list[str] | None = None): return await asyncio.gather( @@ -160,9 +178,11 @@ async def _check(remaining_seconds: float) -> None: @lock_exempt def _count_startable_cells(self) -> int: - if self.args.colocate: - return len(self.server_cells) - return sum(1 for cell in self.server_cells.values() if cell.is_pending_weights_or_serving) + return sum( + 1 + for cell in self.server_cells.values() + if (self.args.colocate and cell.meta.needs_offload) or cell.is_pending_weights_or_serving + ) @property @requires_lock diff --git a/miles/ray/rollout/server_cell.py b/miles/ray/rollout/server_cell.py index d12d6636f26..e598c02d26d 100644 --- a/miles/ray/rollout/server_cell.py +++ b/miles/ray/rollout/server_cell.py @@ -1,6 +1,7 @@ import asyncio import dataclasses import logging +from collections.abc import Callable from dataclasses import dataclass from typing import Any, Literal @@ -18,6 +19,13 @@ StateServing, StateUninitialized, ) +from miles.utils.ft_utils.health_checker import ( + ActiveAndEpoch, + BaseHealthChecker, + NoopHealthChecker, + SimpleHealthChecker, + SimpleHealthCheckerConfig, +) from miles.utils.pydantic_utils import FrozenStrictBaseModel from miles.utils.workers.launch_gate import GATE_PORT_NAME, activate_launch_gate from miles.utils.workers.worker_provider.base import BaseWorkerProvider @@ -46,8 +54,35 @@ class ServerCell: args: Any meta: ServerCellMetadata router_api_client: SGLangRouterApiClient + global_health_checker_activeness: Callable[[], ActiveAndEpoch] = lambda: ActiveAndEpoch(active=True, epoch=0) + _health_checker: BaseHealthChecker = dataclasses.field(init=False) _state: CellState = dataclasses.field(default_factory=StateUninitialized) + def __post_init__(self) -> None: + self._health_checker = create_rollout_cell_health_checker( + args=self.args, + name=f"rollout-cell-{self.meta.cell_id}", + get_api_client=lambda: self.api_client, + get_activeness=self._get_health_checker_active_and_epoch, + ) + self._health_checker.start() + + def _get_health_checker_active_and_epoch(self) -> ActiveAndEpoch: + controller_active_and_epoch = self.global_health_checker_activeness() + cell_active = isinstance(self._state, (StatePendingWeights, StateServing)) + return ActiveAndEpoch( + active=cell_active and controller_active_and_epoch.active, epoch=controller_active_and_epoch.epoch + ) + + def __del__(self) -> None: + assert isinstance(self._state, StateDisposed), ( + f"ServerCell {self.meta.cell_id} was garbage collected without dispose() ({self._state=}); " + "every cell must be disposed so its health checker task is stopped" + ) + + async def cancel_inflight_health_probe(self) -> None: + await self._health_checker.cancel_inflight_probe() + @property def is_uninitialized(self) -> bool: return isinstance(self._state, StateUninitialized) @@ -122,9 +157,7 @@ async def _tick_when_initializing(self) -> None: async def mark_weights_ready(self) -> None: assert isinstance(self._state, StatePendingWeights), f"{self._state=}" - await self._register_with_router(addr_info=self._state.addr_info) - self._mark_serving() async def _register_with_router(self, addr_info: CellAddrInfo) -> None: @@ -136,6 +169,8 @@ async def _register_with_router(self, addr_info: CellAddrInfo) -> None: ) async def dispose(self) -> None: + self._health_checker.stop() + match self._state: case StateServing(): await self._unregister_from_router() @@ -211,3 +246,21 @@ async def check_weights(self, action: str, allow_quant_error: bool, selector: st def compute_nodes_per_engine(*, num_gpus_per_engine: int, num_gpus_per_node: int) -> int: return max(1, num_gpus_per_engine // num_gpus_per_node) + + +def create_rollout_cell_health_checker( + *, + args: Any, + name: str, + get_api_client: Callable[[], SGLangApiClient], + get_activeness: Callable[[], ActiveAndEpoch], +) -> BaseHealthChecker: + if "rollout" not in args.ft_components: + return NoopHealthChecker() + + config = SimpleHealthCheckerConfig.from_args(args, prefix="rollout_health_check") + + async def _check() -> None: + await get_api_client().health_generate(timeout=config.timeout) + + return SimpleHealthChecker(name=name, check_fn=_check, get_activeness=get_activeness, config=config) diff --git a/miles/ray/train/cell_monitor.py b/miles/ray/train/cell_monitor.py index 802be6836de..5006648558d 100644 --- a/miles/ray/train/cell_monitor.py +++ b/miles/ray/train/cell_monitor.py @@ -10,7 +10,7 @@ StateStopped, ) from miles.utils.ft_utils.api_server.models import CellCondition, CellStatus, TriState -from miles.utils.ft_utils.health_checker import ActivenessState, SimpleHealthChecker, SimpleHealthCheckerConfig +from miles.utils.ft_utils.health_checker import ActiveAndEpoch, SimpleHealthChecker, SimpleHealthCheckerConfig if TYPE_CHECKING: from miles.ray.train.cell import RayTrainCell @@ -20,7 +20,7 @@ def create_trainer_cell_health_checker( *, cell: "RayTrainCell", config: SimpleHealthCheckerConfig, - get_activeness: Callable[[], ActivenessState], + get_activeness: Callable[[], ActiveAndEpoch], ) -> SimpleHealthChecker: async def _check() -> None: # Cell health is liveness, not training progress: the heartbeat RPC runs on diff --git a/miles/utils/arguments.py b/miles/utils/arguments.py index a63a60a56ec..70a30eebbce 100644 --- a/miles/utils/arguments.py +++ b/miles/utils/arguments.py @@ -979,6 +979,7 @@ def add_fault_tolerance_arguments(parser): interval_default=30.0, timeout_default=30.0, first_wait_default=0.0, + failure_threshold_default=1, ) parser.add_argument( "--api-server-port", diff --git a/miles/utils/ft_utils/health_checker.py b/miles/utils/ft_utils/health_checker.py index 5c31c234b95..5d919dde95f 100644 --- a/miles/utils/ft_utils/health_checker.py +++ b/miles/utils/ft_utils/health_checker.py @@ -29,6 +29,7 @@ def add_arguments( interval_default: float = 10.0, timeout_default: float = 10.0, first_wait_default: float = 300.0, + failure_threshold_default: int = 3, ) -> None: parser.add_argument( f"--{prefix}-interval", @@ -55,7 +56,7 @@ def add_arguments( parser.add_argument( f"--{prefix}-failure-threshold", type=int, - default=3, + default=failure_threshold_default, help=( f"Number of consecutive failed {prefix} checks before reporting unhealthy. " "Debounces transient RPC blips so a single hiccup does not recycle a live cell." @@ -73,22 +74,22 @@ def from_args(args: object, *, prefix: str) -> SimpleHealthCheckerConfig: ) -class ActivenessState(NamedTuple): +class ActiveAndEpoch(NamedTuple): active: bool epoch: int class ActivenessTracker: def __init__(self, *, active: bool) -> None: - self._state = ActivenessState(active=active, epoch=0) + self._state = ActiveAndEpoch(active=active, epoch=0) - def get(self) -> ActivenessState: + def get(self) -> ActiveAndEpoch: return self._state def bump_active(self, active: bool) -> None: if active == self._state.active: return - self._state = ActivenessState(active=active, epoch=self._state.epoch + 1) + self._state = ActiveAndEpoch(active=active, epoch=self._state.epoch + 1) class BaseHealthChecker(abc.ABC): @@ -102,6 +103,9 @@ def start(self) -> None: ... @abc.abstractmethod def stop(self) -> None: ... + @abc.abstractmethod + async def cancel_inflight_probe(self) -> None: ... + class SimpleHealthChecker(BaseHealthChecker): """Periodic async health checker. Calls *check_fn*; reports result via *on_result*. @@ -115,7 +119,7 @@ def __init__( *, name: str, check_fn: Callable[[], Coroutine[Any, Any, None]], - get_activeness: Callable[[], ActivenessState], + get_activeness: Callable[[], ActiveAndEpoch], on_result: Callable[[bool], None] | None = None, config: SimpleHealthCheckerConfig, clock: Clock | None = None, @@ -128,10 +132,12 @@ def __init__( self._clock = clock or RealClock() self._status = TriState.UNKNOWN - self._activeness = ActivenessState(active=False, epoch=0) + self._active_and_epoch = ActiveAndEpoch(active=False, epoch=0) self._need_first_wait: bool = True self._consecutive_failures: int = 0 self._task: asyncio.Task[None] | None = None + self._probe_task: asyncio.Task[None] | None = None + self._probe_discarded: bool = False @property def status(self) -> TriState: @@ -151,8 +157,29 @@ def stop(self) -> None: log_structured(logger.info, tag="ft", op="health", phase="stop", name=self._name) self._task.cancel() self._task = None + if self._probe_task is not None: + self._probe_task.cancel() + self._probe_task = None self._status = TriState.UNKNOWN + async def cancel_inflight_probe(self) -> None: + probe_task = self._probe_task + if probe_task is None: + return + + log_structured(logger.info, tag="ft", op="health", phase="cancel_probe", name=self._name) + self._probe_discarded = True + probe_task.cancel() + + try: + await probe_task + except asyncio.CancelledError: + pass + except Exception: + log_structured( + logger.error, tag="ft", op="health", phase="cancel_probe_failed", name=self._name, exc_info=True + ) + def _on_paused(self) -> None: log_structured(logger.info, tag="ft", op="health", phase="pause", name=self._name) self._status = TriState.UNKNOWN @@ -165,10 +192,10 @@ def _on_resumed(self) -> None: async def _loop(self) -> None: while True: - activeness = self._get_activeness() - active = activeness.active - if activeness != self._activeness: - self._activeness = activeness + active_and_epoch = self._get_activeness() + active = active_and_epoch.active + if active_and_epoch != self._active_and_epoch: + self._active_and_epoch = active_and_epoch if active: self._on_resumed() else: @@ -188,62 +215,78 @@ async def _loop(self) -> None: continue if active: - success = False - try: - await asyncio.wait_for(self._check_fn(), timeout=self._config.timeout) - success = True - except Exception: - log_structured( - logger.error, tag="ft", op="health", phase="check_failed", name=self._name, exc_info=True - ) - - prev_status = self._status - if success: - self._consecutive_failures = 0 - self._status = TriState.TRUE + success = await self._run_probe() + if self._probe_discarded: + self._probe_discarded = False + log_structured(logger.info, tag="ft", op="health", phase="probe_discarded", name=self._name) else: - self._consecutive_failures += 1 - if self._consecutive_failures >= self._config.failure_threshold: - self._status = TriState.FALSE + self._publish_result(success=success) + + await self._clock.sleep(self._config.interval) + async def _run_probe(self) -> bool: + self._probe_discarded = False + self._probe_task = asyncio.create_task(self._check_fn()) + + try: + await asyncio.wait_for(self._probe_task, timeout=self._config.timeout) + return True + except asyncio.CancelledError: + if not self._probe_discarded: + raise + return False + except Exception: + log_structured(logger.error, tag="ft", op="health", phase="check_failed", name=self._name, exc_info=True) + return False + finally: + self._probe_task = None + + def _publish_result(self, *, success: bool) -> None: + prev_status = self._status + if success: + self._consecutive_failures = 0 + self._status = TriState.TRUE + else: + self._consecutive_failures += 1 + if self._consecutive_failures >= self._config.failure_threshold: + self._status = TriState.FALSE + + log_structured( + logger.info, + tag="ft", + op="health", + phase="poll", + name=self._name, + ok=success, + status=self._status.value, + consecutive_failures=self._consecutive_failures, + ) + + if prev_status != self._status: + log_structured( + logger.info, + tag="ft", + op="health", + phase="status_change", + name=self._name, + from_status=prev_status.value, + to_status=self._status.value, + consecutive_failures=self._consecutive_failures, + ) + + if self._on_result is not None: + try: + self._on_result(success) + except Exception: log_structured( - logger.info, + logger.error, tag="ft", op="health", - phase="poll", + phase="on_result_failed", name=self._name, - ok=success, - status=self._status.value, - consecutive_failures=self._consecutive_failures, + exc_info=True, ) - if prev_status != self._status: - log_structured( - logger.info, - tag="ft", - op="health", - phase="status_change", - name=self._name, - from_status=prev_status.value, - to_status=self._status.value, - consecutive_failures=self._consecutive_failures, - ) - - if self._on_result is not None: - try: - self._on_result(success) - except Exception: - log_structured( - logger.error, - tag="ft", - op="health", - phase="on_result_failed", - name=self._name, - exc_info=True, - ) - - await self._clock.sleep(self._config.interval) - class NoopHealthChecker(BaseHealthChecker): @property @@ -255,3 +298,6 @@ def start(self) -> None: def stop(self) -> None: pass + + async def cancel_inflight_probe(self) -> None: + pass diff --git a/tests/fast/ray/rollout/conftest.py b/tests/fast/ray/rollout/conftest.py index 54b655424ab..0cf761ee4f8 100644 --- a/tests/fast/ray/rollout/conftest.py +++ b/tests/fast/ray/rollout/conftest.py @@ -2,7 +2,8 @@ import textwrap from argparse import ArgumentParser, Namespace -from typing import TYPE_CHECKING, Any +from collections.abc import AsyncIterator +from typing import Any from unittest.mock import MagicMock import pytest @@ -12,21 +13,6 @@ from miles.utils import object_store from miles.utils.types import Sample -if TYPE_CHECKING: - from miles.ray.rollout.server_cell import ServerCell - - -def fake_actor_handle() -> MagicMock: - """MagicMock that passes ``isinstance(x, ray.actor.ActorHandle)``. - - Setting ``_spec_class`` directly (rather than ``spec=...``) keeps - arbitrary-attribute auto-creation working so ``actor.shutdown.remote(...)`` - chains still resolve — ``ActorHandle`` routes its methods via - ``__getattr__`` and they don't show up as class attributes.""" - m = MagicMock() - m._spec_class = ray.actor.ActorHandle - return m - def make_args(**overrides: Any) -> Namespace: """Args namespace covering every field touched by ``miles/ray/rollout/``. @@ -113,8 +99,10 @@ def make_args(**overrides: Any) -> Namespace: offload_rollout=False, use_fault_tolerance=False, ft_components=[], - rollout_health_check_interval=10.0, + rollout_health_check_interval=30.0, rollout_health_check_timeout=30.0, + rollout_health_check_first_wait=0.0, + rollout_health_check_failure_threshold=1, # engine launch command seed=42, fp16=False, @@ -232,6 +220,27 @@ def make_sglang_config_yaml( return "\n".join(lines) + "\n" +# --------------------------- server cell fixtures --------------------------- + +_tracked_server_cells: list[Any] = [] + + +def track_server_cell(cell: Any) -> Any: + """Register a cell for teardown. ``ServerCell.__del__`` asserts that every cell was disposed.""" + _tracked_server_cells.append(cell) + return cell + + +@pytest.fixture +async def dispose_tracked_server_cells() -> AsyncIterator[None]: + """Dispose every cell registered through ``track_server_cell`` during the test.""" + _tracked_server_cells.clear() + yield + for cell in _tracked_server_cells: + await cell.dispose() + _tracked_server_cells.clear() + + # --------------------------- ray fixtures --------------------------- @@ -299,37 +308,6 @@ def dedent(s: str) -> str: return textwrap.dedent(s).lstrip("\n") -def make_dataclass_cells( - *, - num_cells: int = 2, - num_gpus_per_engine: int = 1, - gpu_offset: int = 0, -) -> list[ServerCell]: - """Build configured ``ServerCell``s. Each cell starts unallocated.""" - from miles.ray.rollout.server_cell import ServerCell, ServerCellMetadata - - args = make_args(num_gpus_per_node=8) - return [ - ServerCell( - args=args, - meta=ServerCellMetadata( - worker_type="regular", - cell_id=f"cell-{cell_index}", - num_gpus_per_engine=num_gpus_per_engine, - gpu_offset=gpu_offset + cell_index * min(num_gpus_per_engine, 8), - sglang_overrides={}, - model_idx=0, - group_index=0, - cell_index=cell_index, - needs_offload=False, - model_path=None, - update_weights=False, - ), - ) - for cell_index in range(num_cells) - ] - - def fake_engine(host: str = "10.0.0.1", port_seed: int = 30000) -> MagicMock: """MagicMock that mimics the engine ``CommandActor`` enough for ``addr_allocator``. diff --git a/tests/fast/ray/rollout/real_ray/test_inference_controller.py b/tests/fast/ray/rollout/real_ray/test_inference_controller.py index 55cc92cdc46..a489d9a0b1e 100644 --- a/tests/fast/ray/rollout/real_ray/test_inference_controller.py +++ b/tests/fast/ray/rollout/real_ray/test_inference_controller.py @@ -491,62 +491,24 @@ async def test_recovers_dead_engine_after_rollout_started( @pytest.mark.asyncio class TestRolloutFaultToleranceIsUnsupported: - async def test_health_monitoring_hooks_are_noops_without_fault_tolerance( + async def test_fault_injection_is_skipped_when_fault_tolerance_skips_rollout( self, ray_local_mode, placement_group_factory, tmp_path, patch_low_level, ): - """A plain run never asked for fault tolerance, so the hooks stay out of its way.""" - args = _make_test_args(tmp_path, models=[("actor", True)]) - pg = placement_group_factory(2) - - controller = InferenceController(args, pg) - await controller.init() - - await controller.health_monitoring_pause() - await controller.health_monitoring_resume() - - async def test_health_monitoring_hooks_refuse_to_run_under_fault_tolerance( - self, - ray_local_mode, - placement_group_factory, - tmp_path, - patch_low_level, - ): - """Asking for fault tolerance must fail loudly, not run unmonitored.""" - args = _make_test_args(tmp_path, models=[("actor", True)]) - pg = placement_group_factory(2) - - controller = InferenceController(args, pg) - await controller.init() - controller.args.use_fault_tolerance = True - controller.args.ft_components = ["rollout"] - - with pytest.raises(NotImplementedError): - await controller.health_monitoring_pause() - with pytest.raises(NotImplementedError): - await controller.health_monitoring_resume() - - async def test_health_monitoring_hooks_are_noops_when_fault_tolerance_skips_rollout( - self, - ray_local_mode, - placement_group_factory, - tmp_path, - patch_low_level, - ): - """Fault tolerance limited to training never monitored the engines, so nothing is lost.""" + """Fault tolerance limited to training never reaches the rollout fault injector.""" args = _make_test_args(tmp_path, models=[("actor", True)]) pg = placement_group_factory(2) controller = InferenceController(args, pg) await controller.init() + controller.args.ci_test = True controller.args.use_fault_tolerance = True controller.args.ft_components = ["train"] - await controller.health_monitoring_pause() - await controller.health_monitoring_resume() + await controller.prepare_rollout(2) async def test_fault_injection_refuses_to_run( self, diff --git a/tests/fast/ray/rollout/test_inference_controller.py b/tests/fast/ray/rollout/test_inference_controller.py index caed2ffdaed..a19877e7f08 100644 --- a/tests/fast/ray/rollout/test_inference_controller.py +++ b/tests/fast/ray/rollout/test_inference_controller.py @@ -2,6 +2,7 @@ from argparse import Namespace from pathlib import Path from types import SimpleNamespace +from typing import Any import pytest from tests.fast.ray.rollout.conftest import make_args @@ -12,6 +13,7 @@ from miles.ray.rollout.server_cell import ServerCellMetadata from miles.ray.specs.inference import compute_engine_pool_ids, compute_router_pool_id, specs_inference_engine from miles.utils.context_lock import ContextLock +from miles.utils.ft_utils.health_checker import ActivenessTracker from miles.utils.workers.worker_provider.base import CellInfo, ReconcileFn, StopWatchFn from miles.utils.workers.worker_spec import WorkerMetaContext @@ -58,9 +60,21 @@ def _make_cell_meta(info: CellInfo) -> ServerCellMetadata: class _RecordingServer: - def __init__(self, server_cells: dict | None = None): + def __init__(self, server_cells: dict | None = None, *, model_name: str = "model", update_weights: bool = False): self.server_cells = server_cells or {} + self.update_weights = update_weights + self.model_name = model_name self.calls: list[tuple] = [] + self.api_clients: list = [] + self.engine_gpu_counts: list[int] = [] + self.engine_gpu_offsets: list[int] = [] + + async def offload(self, tags=None): + self.calls.append(("offload",)) + + async def check_weights(self, action, allow_quant_error=False, selector="all", skip_list=None): + self.calls.append(("check_weights", action)) + return [self.model_name] async def add_cell(self, cell_meta: ServerCellMetadata): self.calls.append(("add", cell_meta.cell_id)) @@ -76,12 +90,62 @@ async def wait_expected_num_cells(self) -> None: def _make_controller(servers: dict) -> InferenceController: controller = InferenceController.__new__(InferenceController) - controller.args = SimpleNamespace(debug_train_only=False, use_fault_tolerance=False) + controller.args = SimpleNamespace( + debug_train_only=False, + use_fault_tolerance=False, + ci_test=False, + colocate=False, + offload_rollout_level=["kv_cache", "weight"], + ) controller.servers = servers controller.context_lock = ContextLock("InferenceController") + controller._health_checker_activeness = ActivenessTracker(active=True) return controller +class TestHealthCheckerActiveness: + @pytest.mark.asyncio + async def test_offload_pauses_probing_before_putting_engines_to_sleep(self): + """A slept engine cannot answer /health_generate, so probing must stop first.""" + srv = _RecordingServer() + controller = _make_controller({"default": srv}) + + await controller.offload() + + assert not controller._health_checker_activeness.get().active + assert srv.calls == [("offload",)] + + @pytest.mark.asyncio + async def test_starting_a_weight_update_pauses_probing(self): + """Engines are unusable while their weights are being replaced.""" + controller = _make_controller({"default": _RecordingServer()}) + + info = await controller.start_update_weights() + await controller.end_update_weights(snapshot_cell_id_to_hashes=info.snapshot_cell_id_to_hashes) + + assert not controller._health_checker_activeness.get().active + + @pytest.mark.asyncio + async def test_preparing_a_rollout_resumes_probing(self): + """Probing comes back exactly when the engines start serving traffic again.""" + controller = _make_controller({"default": _RecordingServer()}) + controller._health_checker_activeness.bump_active(False) + + await controller.prepare_rollout(rollout_id=0) + + assert controller._health_checker_activeness.get().active + + @pytest.mark.asyncio + async def test_preparing_an_eval_resumes_probing(self): + """Eval drives the same engines as a rollout does.""" + controller = _make_controller({"default": _RecordingServer()}) + controller._health_checker_activeness.bump_active(False) + + await controller.prepare_eval() + + assert controller._health_checker_activeness.get().active + + class TestReconcile: @pytest.fixture def servers(self) -> dict[str, _RecordingServer]: @@ -194,9 +258,7 @@ async def _stop_watch() -> None: def _patch_init( monkeypatch: pytest.MonkeyPatch, *, provider: _FakeWorkerProvider, servers: dict[str, _RecordingServer] ) -> None: - async def _fake_create_rollout_servers( - args: Namespace, *, context_lock: ContextLock - ) -> dict[str, _RecordingServer]: + async def _fake_create_rollout_servers(args: Namespace, **kwargs: Any) -> dict[str, _RecordingServer]: return servers monkeypatch.setattr(inference_controller_module, "create_rollout_servers", _fake_create_rollout_servers) @@ -208,6 +270,33 @@ async def _fake_create_rollout_servers( ) +class TestGlobalHealthCheckerActiveness: + @pytest.mark.asyncio + async def test_init_hands_the_cells_the_controller_wide_activeness(self, monkeypatch: pytest.MonkeyPatch): + """Without it every cell keeps probing through the weight-update window the controller + just paused, and reports a mid-update engine unhealthy.""" + received: dict[str, Any] = {} + + async def _fake_create_rollout_servers(args: Namespace, **kwargs: Any) -> dict[str, _RecordingServer]: + received.update(kwargs) + return {"default": _RecordingServer()} + + monkeypatch.setattr(inference_controller_module, "create_rollout_servers", _fake_create_rollout_servers) + monkeypatch.setattr( + inference_controller_module, + "RayWorkerProvider", + SimpleNamespace(create=lambda *, pool_ids: _FakeWorkerProvider([]).created_with(pool_ids)), + ) + controller = InferenceController(make_args()) + + await controller.init() + + get_activeness = received["global_health_checker_activeness"] + assert get_activeness().active is True + controller._health_checker_activeness.bump_active(False) + assert get_activeness().active is False + + class TestInitSubscription: @pytest.mark.asyncio async def test_init_watches_exactly_the_engine_specs(self, monkeypatch: pytest.MonkeyPatch): @@ -337,3 +426,59 @@ async def test_a_server_holding_a_foreign_lock_is_rejected(self): with pytest.raises(AssertionError, match="must be called with"): await controller.offload() + + +class TestUpdatableModelSelection: + @staticmethod + def _controller(*servers: _RecordingServer) -> InferenceController: + return _make_controller({srv.model_name: srv for srv in servers}) + + @pytest.mark.asyncio + async def test_only_the_updatable_models_engines_receive_weights(self): + """A frozen reference model handed the trainer's weights stops being the baseline the + KL term is measured against.""" + actor = _RecordingServer(model_name="actor", update_weights=True) + actor.api_clients = ["actor-client"] + ref = _RecordingServer(model_name="ref", update_weights=False) + ref.api_clients = ["ref-client"] + + updatable = await self._controller(actor, ref).start_update_weights() + + assert updatable.rollout_engines == ["actor-client"] + + @pytest.mark.asyncio + async def test_an_inference_only_deployment_updates_nothing(self): + """No model is being trained, so there is no engine to push weights into; returning a + frozen model's engines here would overwrite it.""" + updatable = await self._controller(_RecordingServer(model_name="ref")).start_update_weights() + + assert updatable.rollout_engines == [] + assert updatable.snapshot_cell_id_to_hashes == {} + + @pytest.mark.asyncio + async def test_two_updatable_models_are_refused_by_name(self): + """Picking one arbitrarily would silently train one model and leave the other stale.""" + controller = self._controller( + _RecordingServer(model_name="a", update_weights=True), + _RecordingServer(model_name="b", update_weights=True), + ) + + with pytest.raises(ValueError, match="Multiple servers have update_weights=True"): + await controller.start_update_weights() + + @pytest.mark.asyncio + async def test_the_weight_checker_skips_the_frozen_models(self): + """reset_tensors on a model nobody will rewrite scrambles it for the rest of the run.""" + actor = _RecordingServer(model_name="actor", update_weights=True) + ref = _RecordingServer(model_name="ref", update_weights=False) + + assert await self._controller(actor, ref).check_weights(action="snapshot") == ["actor"] + assert ref.calls == [] + + @pytest.mark.asyncio + async def test_the_weight_checker_is_a_noop_without_an_updatable_model(self): + """Nothing was updated, so there is nothing to compare against.""" + ref = _RecordingServer(model_name="ref") + + assert await self._controller(ref).check_weights(action="compare") == [] + assert ref.calls == [] diff --git a/tests/fast/ray/rollout/test_inference_controller_tick.py b/tests/fast/ray/rollout/test_inference_controller_tick.py index dfe2268b9ad..7125318c95c 100644 --- a/tests/fast/ray/rollout/test_inference_controller_tick.py +++ b/tests/fast/ray/rollout/test_inference_controller_tick.py @@ -30,6 +30,11 @@ async def tick(self) -> None: class _StubServer: def __init__(self, server_cells: dict): self.server_cells = server_cells + self.dispose_count = 0 + + async def dispose(self) -> None: + self.dispose_count += 1 + self.server_cells.clear() def _make_controller(servers: dict) -> InferenceController: @@ -120,6 +125,15 @@ async def test_dispose_stops_the_ticker(self): assert cell.tick_count == ticks_after_dispose + async def test_dispose_tears_down_every_server_so_no_cell_keeps_probing(self): + """A cell health checker keeps calling /health_generate unless dispose reaches its cell.""" + first, second = _StubServer({"a": _RecordingCell()}), _StubServer({"b": _RecordingCell()}) + controller = _make_controller({"default": first, "frozen": second}) + + await controller.dispose() + + assert (first.dispose_count, second.dispose_count) == (1, 1) + async def test_dispose_without_a_running_ticker_is_harmless(self): """debug_train_only never starts the ticker, and teardown still has to work.""" controller = _make_controller({}) diff --git a/tests/fast/ray/rollout/test_rollout_server.py b/tests/fast/ray/rollout/test_rollout_server.py index f2739008843..65794e63eaa 100644 --- a/tests/fast/ray/rollout/test_rollout_server.py +++ b/tests/fast/ray/rollout/test_rollout_server.py @@ -1,20 +1,23 @@ from __future__ import annotations +import asyncio from types import SimpleNamespace -from typing import Any -from unittest.mock import patch import pytest -from tests.fast.ray.rollout.conftest import make_args, make_dataclass_cells +from tests.fast.ray.rollout.conftest import make_args from miles.backends.sglang_utils.sglang_config import ( _compute_megatron_num_gpus, _compute_rollout_offset, resolve_sglang_config, ) -from miles.ray.rollout.rollout_server import RolloutServer +from miles.ray.rollout import rollout_server as rollout_server_module +from miles.ray.rollout.cell_state import CellAddrInfo, StateServing +from miles.ray.rollout.rollout_server import RolloutServer, create_rollout_servers from miles.ray.rollout.server_cell import ServerCell, ServerCellMetadata from miles.utils.context_lock import ContextLock +from miles.utils.ft_utils.health_checker import ActiveAndEpoch +from miles.utils.workers.worker_spec import HostAndPort class TestRolloutServerPureFunctions: @@ -222,47 +225,59 @@ def test_compute_megatron_num_gpus_zero_when_debug_rollout_only(self): assert _compute_megatron_num_gpus(args) == 0 -@pytest.mark.skip( - reason="TODO: rebuild against the meta/router_api_client ServerCell; make_dataclass_cells and " - "_mark_allocated_uninitialized/_mark_addressing target the removed constructor and state API" -) -class TestRolloutServerCrossCellProperties: - def test_api_clients_expose_one_client_per_cell(self): - """Each cell is addressed through its primary (node-0) endpoint.""" - cells = make_dataclass_cells(num_cells=2, gpu_offset=0) + make_dataclass_cells(num_cells=2, gpu_offset=2) - for index, cell in enumerate(cells): - cell._mark_allocated_uninitialized() - cell._mark_addressing(server_url=f"http://10.0.0.{index + 1}:30000") - srv = RolloutServer( - server_cells={f"cell-{i}": cell for i, cell in enumerate(cells)}, - args=SimpleNamespace(), - context_lock=_make_lock(), - ) - assert [client.server_url for client in srv.api_clients] == [ - f"http://10.0.0.{index + 1}:30000" for index in range(4) - ] +class TestCreateRolloutServersWiring: + _CONFIG_YAML = ( + "sglang:\n" + " - name: actor\n" + " server_groups:\n" + " - worker_type: regular\n" + " num_gpus: 8\n" + " num_gpus_per_engine: 4\n" + " - name: ref\n" + " model_path: /fake/ref-model\n" + " update_weights: false\n" + " server_groups:\n" + " - worker_type: regular\n" + " num_gpus: 4\n" + " num_gpus_per_engine: 4\n" + ) - def test_engine_gpu_counts_parallel_to_engines(self): - cells = make_dataclass_cells(num_cells=2, num_gpus_per_engine=1) + make_dataclass_cells( - num_cells=2, num_gpus_per_engine=2 - ) - srv = RolloutServer( - server_cells={f"cell-{i}": cell for i, cell in enumerate(cells)}, - args=SimpleNamespace(), - context_lock=_make_lock(), - ) - assert srv.engine_gpu_counts == [1, 1, 2, 2] + @pytest.mark.asyncio + async def test_create_rollout_servers_wires_each_model_and_router_completely(self, tmp_path, monkeypatch): + """Every model gets its own router, the first one is published on the legacy args, and the injected lock, activeness getter and update_weights all survive.""" - def test_engine_gpu_offsets_consistent_across_cells(self): - cells = make_dataclass_cells(num_cells=2, num_gpus_per_engine=1, gpu_offset=0) + make_dataclass_cells( - num_cells=2, num_gpus_per_engine=2, gpu_offset=4 - ) - srv = RolloutServer( - server_cells={f"cell-{i}": cell for i, cell in enumerate(cells)}, - args=SimpleNamespace(), - context_lock=_make_lock(), + async def _wait_router_ready(model_idx: int) -> HostAndPort: + return HostAndPort(host=f"10.0.0.{model_idx + 1}", port=20000 + model_idx) + + monkeypatch.setattr(rollout_server_module, "wait_router_ready", _wait_router_ready) + + cfg_path = tmp_path / "cfg.yaml" + cfg_path.write_text(self._CONFIG_YAML) + args = make_args(sglang_config=str(cfg_path), rollout_num_gpus=12, debug_rollout_only=True) + lock = ContextLock("InferenceControllerUnderTest") + active_and_epoch = ActiveAndEpoch(active=False, epoch=7) + + def _get_active_and_epoch() -> ActiveAndEpoch: + return active_and_epoch + + servers = await create_rollout_servers( + args, + context_lock=lock, + global_health_checker_activeness=_get_active_and_epoch, ) - assert srv.engine_gpu_offsets == [0, 1, 4, 6] + + assert {name: (srv.router_ip, srv.router_port) for name, srv in servers.items()} == { + "actor": ("10.0.0.1", 20000), + "ref": ("10.0.0.2", 20001), + } + assert (args.sglang_router_ip, args.sglang_router_port) == ("10.0.0.1", 20000) + assert args.sglang_model_routers == {"actor": ("10.0.0.1", 20000), "ref": ("10.0.0.2", 20001)} + assert [srv.model_name for srv in servers.values()] == ["actor", "ref"] + assert (servers["actor"].update_weights, servers["ref"].update_weights) == (True, False) + active_and_epoch = ActiveAndEpoch(active=True, epoch=8) + for srv in servers.values(): + assert srv.context_lock is lock + assert srv.global_health_checker_activeness() == ActiveAndEpoch(active=True, epoch=8) class TestEngineListOrdering: @@ -301,27 +316,100 @@ def _make_meta( ) @pytest.mark.asyncio - async def test_a_failed_add_leaves_no_bookkeeping_so_the_next_reconcile_retries(self, monkeypatch): - """A cell whose startup fails must not be committed, otherwise the hash no-op blocks any retry.""" - srv = RolloutServer(server_cells={}, args=SimpleNamespace(colocate=False), context_lock=_make_lock()) + async def test_a_failed_add_still_tracks_the_cell_so_nothing_leaks(self, monkeypatch): + """The cell is committed before init runs, so a failing init cannot orphan its health checker task.""" + srv = RolloutServer( + server_cells={}, args=SimpleNamespace(colocate=False, ft_components=[]), context_lock=_make_lock() + ) monkeypatch.setattr(ServerCell, "init", _raise_async) async with srv.context_lock: with pytest.raises(RuntimeError, match="injected init failure"): await srv.add_cell(self._make_meta()) + assert list(srv.server_cells) == ["inference-engine-0-0-0"] + await srv.dispose() + + @pytest.mark.asyncio + async def test_disposing_the_server_removes_every_cell_it_tracks(self, monkeypatch): + """Controller teardown must reach each cell so its health checker task stops with it.""" + srv = RolloutServer( + server_cells={}, args=SimpleNamespace(colocate=True, ft_components=[]), context_lock=_make_lock() + ) + monkeypatch.setattr(ServerCell, "init", _noop_async) + + async with srv.context_lock: + await srv.add_cell(self._make_meta()) + await srv.dispose() + assert srv.server_cells == {} @pytest.mark.asyncio async def test_a_successful_add_commits_the_cell(self, monkeypatch): """After the failure is gone the same cell id can be added normally.""" - srv = RolloutServer(server_cells={}, args=SimpleNamespace(colocate=False), context_lock=_make_lock()) + srv = RolloutServer( + server_cells={}, args=SimpleNamespace(colocate=False, ft_components=[]), context_lock=_make_lock() + ) monkeypatch.setattr(ServerCell, "init", _noop_async) async with srv.context_lock: await srv.add_cell(self._make_meta()) - assert list(srv.server_cells) == ["inference-engine-0-0-0"] + assert list(srv.server_cells) == ["inference-engine-0-0-0"] + await srv.dispose() + + +class TestDuplicateCellId: + @pytest.mark.asyncio + async def test_adding_a_duplicate_cell_id_preserves_the_original_cell(self, monkeypatch): + """Overwriting the entry would drop the first cell's health checker task and router registration on the floor.""" + srv = RolloutServer( + server_cells={}, args=SimpleNamespace(colocate=True, ft_components=[]), context_lock=_make_lock() + ) + monkeypatch.setattr(ServerCell, "init", _noop_async) + + async with srv.context_lock: + await srv.add_cell(TestAddCellRollback()._make_meta()) + original = srv.server_cells["inference-engine-0-0-0"] + + with pytest.raises(AssertionError): + await srv.add_cell(TestAddCellRollback()._make_meta()) + + assert srv.server_cells["inference-engine-0-0-0"] is original + await srv.dispose() + + +class TestRemoveCellDisposal: + @pytest.mark.asyncio + async def test_remove_cell_disposes_it_before_forgetting_it(self, monkeypatch): + """A forgotten-but-undisposed cell keeps its router registration and health checker task alive forever.""" + dispose_started = asyncio.Event() + allow_dispose = asyncio.Event() + real_dispose = ServerCell.dispose + + async def _blocking_dispose(cell: ServerCell) -> None: + dispose_started.set() + await allow_dispose.wait() + await real_dispose(cell) + + srv = RolloutServer( + server_cells={}, args=SimpleNamespace(colocate=True, ft_components=[]), context_lock=_make_lock() + ) + monkeypatch.setattr(ServerCell, "init", _noop_async) + monkeypatch.setattr(ServerCell, "dispose", _blocking_dispose) + + async with srv.context_lock: + await srv.add_cell(TestAddCellRollback()._make_meta()) + remove_task = asyncio.create_task(srv.remove_cell("inference-engine-0-0-0")) + await asyncio.wait_for(dispose_started.wait(), timeout=1) + + assert list(srv.server_cells) == ["inference-engine-0-0-0"] + assert not remove_task.done() + + allow_dispose.set() + await asyncio.wait_for(remove_task, timeout=1) + + assert srv.server_cells == {} class TestAddCellInitTiming: @@ -333,11 +421,14 @@ async def test_a_disaggregated_cell_is_initialized_as_soon_as_it_appears(self, m async def _record(self) -> None: initialized.append(self.meta.cell_id) - srv = RolloutServer(server_cells={}, args=SimpleNamespace(colocate=False), context_lock=_make_lock()) + srv = RolloutServer( + server_cells={}, args=SimpleNamespace(colocate=False, ft_components=[]), context_lock=_make_lock() + ) monkeypatch.setattr(ServerCell, "init", _record) async with srv.context_lock: await srv.add_cell(TestAddCellRollback()._make_meta()) + await srv.dispose() assert initialized == ["inference-engine-0-0-0"] @@ -349,95 +440,125 @@ async def test_a_colocated_cell_is_only_tracked_until_the_weight_update_window(s async def _record(self) -> None: initialized.append(self.meta.cell_id) - srv = RolloutServer(server_cells={}, args=SimpleNamespace(colocate=True), context_lock=_make_lock()) + srv = RolloutServer( + server_cells={}, args=SimpleNamespace(colocate=True, ft_components=[]), context_lock=_make_lock() + ) monkeypatch.setattr(ServerCell, "init", _record) async with srv.context_lock: await srv.add_cell(TestAddCellRollback()._make_meta(needs_offload=True)) - assert initialized == [] - assert list(srv.server_cells) == ["inference-engine-0-0-0"] + assert initialized == [] + assert list(srv.server_cells) == ["inference-engine-0-0-0"] + await srv.dispose() + @pytest.mark.asyncio + async def test_a_non_offloaded_cell_is_initialized_immediately_even_under_colocate(self, monkeypatch): + """A dedicated eval cell does not share GPUs with the trainer, so it must not wait for the first weight update.""" + initialized: list[str] = [] -def _make_lock() -> ContextLock: - return ContextLock("InferenceController") + async def _record(self) -> None: + initialized.append(self.meta.cell_id) + srv = RolloutServer( + server_cells={}, args=SimpleNamespace(colocate=True, ft_components=[]), context_lock=_make_lock() + ) + monkeypatch.setattr(ServerCell, "init", _record) + meta_builder = TestAddCellRollback() -async def _raise_async(self): - raise RuntimeError("injected init failure") + async with srv.context_lock: + await srv.add_cell(meta_builder._make_meta(needs_offload=True)) + assert initialized == [] -async def _noop_async(self): - return None + await srv.add_cell(meta_builder._make_meta(needs_offload=False, cell_id="eval-engine-0-0-0")) + assert initialized == ["eval-engine-0-0-0"] + await srv.dispose() -@pytest.mark.skip( - reason="TODO: rebuild against the meta/router_api_client ServerCell; _make_started_server still drives the " - "removed AddrInfo/_mark_addressing/_mark_alive state API and the removed constructors" -) -class TestRemoveCell: + +class TestDeferredInitMatchesTheStartupBarrier: + @pytest.mark.parametrize( + ("colocate", "needs_offload", "deferred"), + [(False, False, False), (False, True, False), (True, False, False), (True, True, True)], + ids=["disaggregated_plain", "disaggregated_offloaded", "colocated_eval", "colocated_shared"], + ) @pytest.mark.asyncio - async def test_remove_cell_detaches_the_cell_from_every_server_view(self): - """A removed cell is gone from server_cells and from every view derived from it.""" - events: list[dict[str, Any]] = [] - srv = _make_started_server(num_cells=2) + async def test_only_a_cell_sharing_the_trainer_gpus_is_deferred_and_counted_before_it_is_up( + self, monkeypatch, colocate: bool, needs_offload: bool, deferred: bool + ): + """Whichever cells skip startup init are exactly the ones the barrier may pass without an address.""" + initialized: list[str] = [] - with _with_recording_router(events): - await srv.remove_cell("default-0") + async def _record(self) -> None: + initialized.append(self.meta.cell_id) - assert "default-0" not in srv.server_cells - assert list(srv.server_cells) == ["default-1"] - assert [client.server_url for client in srv.api_clients] == ["http://10.0.0.2:30001"] - assert srv.engine_gpu_counts == [1] - assert srv.engine_gpu_offsets == [1] + srv = RolloutServer( + server_cells={}, args=SimpleNamespace(colocate=colocate, ft_components=[]), context_lock=_make_lock() + ) + monkeypatch.setattr(ServerCell, "init", _record) - @pytest.mark.asyncio - async def test_remove_cell_unregisters_from_the_router_before_dropping_the_cell(self): - """Dropping the cell without unregistering would leave the router routing to a dead worker.""" - events: list[dict[str, Any]] = [] - srv = _make_started_server(num_cells=2) + async with srv.context_lock: + await srv.add_cell(TestAddCellRollback()._make_meta(needs_offload=needs_offload)) - with _with_recording_router(events): - await srv.remove_cell("default-0") + assert (initialized == []) is deferred + assert srv._count_startable_cells() == (1 if deferred else 0) + await srv.dispose() - assert events == [{"call": "remove_worker", "worker_url": "http://10.0.0.1:30000", "use_legacy_api": False}] +async def _make_serving_server(monkeypatch, *, num_cells: int) -> RolloutServer: + """A server whose cells are all registered and serving, without dialling anything.""" + router = _RecordingRouterApiClient() + monkeypatch.setattr(RolloutServer, "_router_api_client", property(lambda self: router)) -def _make_started_server(*, num_cells: int) -> RolloutServer: - args = make_args(num_gpus_per_node=8) - srv = RolloutServer(server_cells={}, args=args) - for cell_index in range(num_cells): - meta = ServerCellMetadata( - model_id="default", - worker_type="regular", - cell_id=f"default-{cell_index}", - num_gpus_per_engine=1, - gpu_offset=cell_index, - sglang_api_key=None, - worker_name=f"default-{cell_index}-0", - needs_offload=False, - update_weights=True, - workers_hash=f"pseudo-hash-{cell_index}", - ) - cell = ServerCell(args=args, meta=meta) - cell._mark_addressing(AddrInfo(server_url=f"http://10.0.0.{cell_index + 1}:3000{cell_index}")) # noqa: F821 - cell._mark_alive() - srv.server_cells[meta.cell_id] = cell + srv = RolloutServer( + server_cells={}, + args=SimpleNamespace(colocate=True, ft_components=[], use_miles_router=False), + context_lock=_make_lock(), + ) + async with srv.context_lock: + for cell_index in range(num_cells): + await srv.add_cell( + ServerCellMetadata( + model_id="default", + worker_type="regular", + cell_id=f"default-{cell_index}", + num_gpus_per_engine=1, + gpu_offset=cell_index, + sglang_api_key=None, + worker_name=f"default-{cell_index}-0", + needs_offload=True, + update_weights=True, + workers_hash=f"pseudo-hash-{cell_index}", + ) + ) + cell = srv.server_cells[f"default-{cell_index}"] + addr_info = CellAddrInfo( + server_url=f"http://10.0.0.{cell_index + 1}:3000{cell_index}", bootstrap_port=None, gate_url=None + ) + await cell._register_with_router(addr_info=addr_info) + cell._state = StateServing(addr_info=addr_info) return srv -def _with_recording_router(events: list[dict[str, Any]]) -> Any: - return patch.object( - RolloutServer, "_router_api_client", property(lambda self: _RecordingRouterApiClient(events=events)) - ) +class _RecordingRouterApiClient: + def __init__(self) -> None: + self.calls: list[tuple[str, dict]] = [] + async def add_worker(self, **kwargs) -> None: + self.calls.append(("add_worker", kwargs)) -class _RecordingRouterApiClient: - def __init__(self, *, events: list[dict[str, Any]]) -> None: - self._events = events + async def remove_worker(self, **kwargs) -> None: + self.calls.append(("remove_worker", kwargs)) + + +def _make_lock() -> ContextLock: + return ContextLock("InferenceController") - async def add_worker(self, **kwargs: Any) -> None: - self._events.append({"call": "add_worker", **kwargs}) - async def remove_worker(self, **kwargs: Any) -> None: - self._events.append({"call": "remove_worker", **kwargs}) +async def _raise_async(self): + raise RuntimeError("injected init failure") + + +async def _noop_async(self): + return None diff --git a/tests/fast/ray/rollout/test_rollout_server_expected_num_cells.py b/tests/fast/ray/rollout/test_rollout_server_expected_num_cells.py index 29de9b61b33..3ee42e9d3ba 100644 --- a/tests/fast/ray/rollout/test_rollout_server_expected_num_cells.py +++ b/tests/fast/ray/rollout/test_rollout_server_expected_num_cells.py @@ -9,6 +9,8 @@ from miles.ray.rollout import rollout_server as rollout_server_module from miles.ray.rollout.rollout_server import create_rollout_servers from miles.ray.specs.inference import compute_engine_pool_id, specs_inference_engine +from miles.utils.context_lock import ContextLock +from miles.utils.ft_utils.health_checker import ActiveAndEpoch from miles.utils.workers.worker_spec import HostAndPort _CONFIG_SINGLE_GROUP: list[dict] = [ @@ -107,6 +109,15 @@ async def _wait_router_ready(model_idx: int) -> HostAndPort: monkeypatch.setattr(rollout_server_module, "wait_router_ready", _wait_router_ready) +async def _create_servers(args) -> dict: + """create_rollout_servers takes the controller's lock and activeness getter in production.""" + return await create_rollout_servers( + args, + context_lock=ContextLock("test"), + global_health_checker_activeness=lambda: ActiveAndEpoch(active=True, epoch=0), + ) + + class TestExpectedNumCellsMatchesTheEngineSpecs: @pytest.mark.parametrize( "models", @@ -126,7 +137,7 @@ async def test_the_startup_barrier_expects_exactly_the_cells_the_specs_launch( args = _make_args_with_config(models=models, tmp_path=tmp_path) expected_per_model_idx = _expected_num_cells_from_specs(args) - servers = await create_rollout_servers(args) + servers = await _create_servers(args) actual_per_model_idx = { model_idx: servers[model["name"]].expected_num_cells for model_idx, model in enumerate(models) @@ -137,7 +148,7 @@ async def test_a_placeholder_group_contributes_no_cell_to_the_barrier(self, tmp_ """Placeholder groups only reserve GPU slots, so counting them would make the barrier unreachable.""" args = _make_args_with_config(models=_CONFIG_WITH_PLACEHOLDER, tmp_path=tmp_path) - servers = await create_rollout_servers(args) + servers = await _create_servers(args) assert servers["actor"].expected_num_cells == 2 @@ -145,7 +156,7 @@ async def test_every_model_gets_its_own_barrier_target(self, tmp_path: Path) -> """Sharing one pool size across models would block the small model behind the big one.""" args = _make_args_with_config(models=_CONFIG_MULTI_MODEL, tmp_path=tmp_path) - servers = await create_rollout_servers(args) + servers = await _create_servers(args) assert servers["actor"].expected_num_cells == 4 assert servers["ref"].expected_num_cells == 1 @@ -155,3 +166,22 @@ class TestEngineSpecNamingUsedByTheCrossCheck: def test_pool_names_carry_the_model_index_the_cross_check_parses(self) -> None: """The cross-check maps specs back to models by name, so that encoding must stay stable.""" assert compute_engine_pool_id(model_idx=3, group_index=7) == "inference-engine-3-7" + + +class TestRouterFlagsAtStartup: + async def test_a_pinned_router_port_is_accepted(self, tmp_path: Path) -> None: + """Three examples pass --sglang-router-port so a firewall rule can name the port in + advance, and the spec layer pins it; rejecting it here fails those launches outright.""" + args = _make_args_with_config(models=_CONFIG_SINGLE_GROUP, tmp_path=tmp_path) + args.sglang_router_port = 31000 + + assert await _create_servers(args) + + async def test_an_external_router_ip_is_still_rejected(self, tmp_path: Path) -> None: + """Attaching to a router miles did not start is not supported yet, and silently starting + one anyway would put two routers in front of the same engines.""" + args = _make_args_with_config(models=_CONFIG_SINGLE_GROUP, tmp_path=tmp_path) + args.sglang_router_ip = "10.0.0.9" + + with pytest.raises(AssertionError, match="external router mode was removed"): + await _create_servers(args) diff --git a/tests/fast/ray/rollout/test_rollout_server_health_checking.py b/tests/fast/ray/rollout/test_rollout_server_health_checking.py new file mode 100644 index 00000000000..970d0579404 --- /dev/null +++ b/tests/fast/ray/rollout/test_rollout_server_health_checking.py @@ -0,0 +1,183 @@ +from types import SimpleNamespace + +import pytest +from tests.fast.ray.rollout.conftest import make_args, track_server_cell + +from miles.ray.rollout import server_cell as server_cell_module +from miles.ray.rollout.cell_state import CellAddrInfo +from miles.ray.rollout.rollout_server import RolloutServer +from miles.ray.rollout.server_cell import ServerCell, ServerCellMetadata +from miles.utils.context_lock import ContextLock +from miles.utils.ft_utils.health_checker import ActiveAndEpoch, NoopHealthChecker, SimpleHealthChecker + +pytestmark = pytest.mark.usefixtures("dispose_tracked_server_cells") + +_ADDR_INFO = CellAddrInfo( + server_url="http://10.0.0.1:30000", + bootstrap_port=None, + gate_url="http://10.0.0.1:13000", +) + + +def _make_meta(cell_id: str = "cell-0", **overrides) -> ServerCellMetadata: + return ServerCellMetadata( + **{ + "model_id": "default", + "worker_type": "regular", + "cell_id": cell_id, + "num_gpus_per_engine": 1, + "gpu_offset": 0, + "sglang_api_key": None, + "worker_name": f"{cell_id}-0", + "needs_offload": False, + "update_weights": True, + "workers_hash": "pseudo-hash-0", + **overrides, + } + ) + + +def _make_server(*, ft_components=("rollout",), **overrides) -> RolloutServer: + return RolloutServer( + server_cells={}, + args=make_args(colocate=True, ft_components=list(ft_components)), + context_lock=ContextLock("InferenceController"), + **overrides, + ) + + +def _make_cell(*, global_activeness=None, router=None, **meta_overrides) -> ServerCell: + return track_server_cell( + ServerCell( + args=make_args(ft_components=["rollout"]), + meta=_make_meta(**meta_overrides), + router_api_client=router or SimpleNamespace(), + global_health_checker_activeness=global_activeness or (lambda: ActiveAndEpoch(active=True, epoch=0)), + ) + ) + + +def _stub_network(monkeypatch, *, ready: bool = True) -> None: + async def _activate(gate_url: str) -> None: + pass + + async def _probe(server_url: str, api_key, timeout: float = 5.0) -> bool: + return ready + + async def _compute_addr_info(self) -> CellAddrInfo: + return _ADDR_INFO + + monkeypatch.setattr(server_cell_module, "activate_launch_gate", _activate) + monkeypatch.setattr(server_cell_module, "probe_server_healthy", _probe) + monkeypatch.setattr(ServerCell, "_compute_addr_info", _compute_addr_info) + + +async def _noop_add_worker(**kwargs) -> None: + pass + + +class TestHealthCheckerActiveness: + async def test_a_gated_cell_is_not_probed(self): + """A colocated cell waits gated for the next window; its port is not even listening.""" + cell = _make_cell() + + assert not cell._get_health_checker_active_and_epoch().active + + async def test_a_booting_cell_is_not_probed(self, monkeypatch): + """Its engine has not answered yet, so a probe would count a false failure.""" + _stub_network(monkeypatch, ready=False) + cell = _make_cell() + + await cell.init() + + assert not cell._get_health_checker_active_and_epoch().active + + async def test_a_cell_holding_stale_weights_is_probed(self, monkeypatch): + """It answers requests with stale weights, so a crash there is a real failure.""" + _stub_network(monkeypatch) + cell = _make_cell() + + await cell.init() + await cell.tick() + + assert cell._get_health_checker_active_and_epoch().active + + async def test_a_cell_that_skips_pending_weights_is_probed(self, monkeypatch): + """A frozen model goes straight to serving, and an unwatched serving cell is a hole in FT.""" + _stub_network(monkeypatch) + cell = _make_cell(router=SimpleNamespace(add_worker=_noop_add_worker), update_weights=False) + + await cell.init() + await cell.tick() + + assert cell._get_health_checker_active_and_epoch().active + + async def test_a_disposed_cell_is_not_probed(self, monkeypatch): + """Nothing is left to answer once the cell has been torn down.""" + _stub_network(monkeypatch) + cell = _make_cell() + + await cell.init() + await cell.tick() + await cell.dispose() + + assert not cell._get_health_checker_active_and_epoch().active + + async def test_a_forced_pause_wins_over_the_cell_state(self, monkeypatch): + """Engines are unusable while offloaded or mid weight update, whatever state they are in.""" + _stub_network(monkeypatch) + active = {"value": True} + cell = _make_cell(global_activeness=lambda: ActiveAndEpoch(active=active["value"], epoch=0)) + + await cell.init() + await cell.tick() + assert cell._get_health_checker_active_and_epoch().active + + active["value"] = False + + assert not cell._get_health_checker_active_and_epoch().active + + +class TestAddCellHealthChecker: + async def test_a_new_cell_starts_its_checker(self): + """Activeness is pulled per loop, so starting unconditionally is safe and needs no replay.""" + srv = _make_server() + + async with srv.context_lock: + await srv.add_cell(_make_meta(needs_offload=True)) + + checker = srv.server_cells["cell-0"]._health_checker + assert isinstance(checker, SimpleHealthChecker) + assert checker._task is not None + await srv.dispose() + + async def test_a_cell_added_mid_window_does_not_probe(self): + """The window releases the lock, so reconcile can add a cell while probing is paused.""" + srv = _make_server(global_health_checker_activeness=lambda: ActiveAndEpoch(active=False, epoch=0)) + + async with srv.context_lock: + await srv.add_cell(_make_meta(needs_offload=True)) + + assert not srv.server_cells["cell-0"]._get_health_checker_active_and_epoch().active + await srv.dispose() + + async def test_no_checker_is_created_without_rollout_fault_tolerance(self): + """Without rollout FT nothing consumes the health status, so nothing probes.""" + srv = _make_server(ft_components=()) + + async with srv.context_lock: + await srv.add_cell(_make_meta(needs_offload=True)) + + assert isinstance(srv.server_cells["cell-0"]._health_checker, NoopHealthChecker) + await srv.dispose() + + +class TestDispose: + async def test_dispose_stops_the_checker_task(self): + """An inactive checker still polls its predicate forever unless the task is cancelled.""" + cell = _make_cell() + assert cell._health_checker._task is not None + + await cell.dispose() + + assert cell._health_checker._task is None diff --git a/tests/fast/ray/rollout/test_rollout_server_locking.py b/tests/fast/ray/rollout/test_rollout_server_locking.py index 8580d1856e7..f0ffb7ee9bf 100644 --- a/tests/fast/ray/rollout/test_rollout_server_locking.py +++ b/tests/fast/ray/rollout/test_rollout_server_locking.py @@ -76,6 +76,6 @@ async def test_cells_can_still_be_added_while_it_polls(self): await asyncio.sleep(0) async with srv.context_lock: - srv.server_cells["inference-engine-0-0-0"] = SimpleNamespace() + srv.server_cells["inference-engine-0-0-0"] = SimpleNamespace(meta=SimpleNamespace(needs_offload=True)) await asyncio.wait_for(waiter, timeout=5) diff --git a/tests/fast/ray/rollout/test_rollout_server_wait_cells.py b/tests/fast/ray/rollout/test_rollout_server_wait_cells.py index 0e8ccf605b4..fc62a5f6041 100644 --- a/tests/fast/ray/rollout/test_rollout_server_wait_cells.py +++ b/tests/fast/ray/rollout/test_rollout_server_wait_cells.py @@ -17,8 +17,9 @@ def fast_polling(monkeypatch): class _FakeCell: - def __init__(self, *, ready: bool = False): + def __init__(self, *, ready: bool = False, needs_offload: bool = True): self.ready = ready + self.meta = SimpleNamespace(needs_offload=needs_offload) @property def is_pending_weights_or_serving(self) -> bool: @@ -76,6 +77,47 @@ async def test_it_returns_once_every_engine_is_up(self): await asyncio.wait_for(task, timeout=5) +class TestWaitExpectedNumCellsWithADedicatedEvalFleet: + async def test_a_cell_outside_the_trainer_gpus_still_has_to_come_up(self): + """A colocated run whose eval cells count as ready on arrival snapshots api clients that have no address yet.""" + srv = _make_server( + colocate=True, expected_num_cells=1, cells={"eval": _FakeCell(ready=False, needs_offload=False)} + ) + + task = asyncio.create_task(srv.wait_expected_num_cells()) + await asyncio.sleep(0) + + assert not task.done() + task.cancel() + + async def test_it_returns_once_the_eval_cell_is_up(self): + """The barrier is what makes the eval fleet see engines that are actually addressable.""" + cell = _FakeCell(ready=False, needs_offload=False) + srv = _make_server(colocate=True, expected_num_cells=1, cells={"eval": cell}) + + task = asyncio.create_task(srv.wait_expected_num_cells()) + await asyncio.sleep(0) + cell.ready = True + + await asyncio.wait_for(task, timeout=5) + + async def test_a_mixed_pool_waits_for_the_eval_cell_but_not_for_the_deferred_ones(self): + """One colocated run holds both kinds of cell, so the wait cannot be decided per run.""" + eval_cell = _FakeCell(ready=False, needs_offload=False) + srv = _make_server( + colocate=True, + expected_num_cells=2, + cells={"shared": _FakeCell(ready=False, needs_offload=True), "eval": eval_cell}, + ) + + task = asyncio.create_task(srv.wait_expected_num_cells()) + await asyncio.sleep(0) + assert not task.done() + + eval_cell.ready = True + await asyncio.wait_for(task, timeout=5) + + class TestWaitExpectedNumCellsEdges: async def test_a_model_without_cells_does_not_wait(self): """A server expecting nothing must not hold up startup.""" diff --git a/tests/fast/ray/rollout/test_server_cell_addressing.py b/tests/fast/ray/rollout/test_server_cell_addressing.py index fd40bdbf725..0a1482bae06 100644 --- a/tests/fast/ray/rollout/test_server_cell_addressing.py +++ b/tests/fast/ray/rollout/test_server_cell_addressing.py @@ -1,12 +1,14 @@ from __future__ import annotations import pytest -from tests.fast.ray.rollout.conftest import make_args +from tests.fast.ray.rollout.conftest import make_args, track_server_cell from miles.ray.rollout import server_cell as server_cell_module from miles.ray.rollout.server_cell import ServerCell, ServerCellMetadata from miles.utils.workers.worker_spec import HostAndPort +pytestmark = pytest.mark.usefixtures("dispose_tracked_server_cells") + def _make_meta(**overrides) -> ServerCellMetadata: return ServerCellMetadata( @@ -49,7 +51,7 @@ def _install(addrs: dict[str, HostAndPort]) -> _StubProvider: def _make_cell(**meta_overrides) -> ServerCell: - return ServerCell(args=make_args(), meta=_make_meta(**meta_overrides), router_api_client=None) + return track_server_cell(ServerCell(args=make_args(), meta=_make_meta(**meta_overrides), router_api_client=None)) class TestComputeAddrInfo: diff --git a/tests/fast/ray/rollout/test_server_cell_health_checker.py b/tests/fast/ray/rollout/test_server_cell_health_checker.py new file mode 100644 index 00000000000..a28cadf795a --- /dev/null +++ b/tests/fast/ray/rollout/test_server_cell_health_checker.py @@ -0,0 +1,287 @@ +from __future__ import annotations + +import asyncio +from typing import Any +from unittest.mock import MagicMock + +import pytest +from tests.fast.ray.rollout.conftest import make_args, track_server_cell + +from miles.ray.rollout import server_cell as server_cell_module +from miles.ray.rollout.cell_state import CellAddrInfo, StatePendingWeights, StateServing, StateUninitialized +from miles.ray.rollout.inference_controller import InferenceController +from miles.ray.rollout.rollout_server import RolloutServer +from miles.ray.rollout.server_cell import ServerCell, ServerCellMetadata +from miles.utils.context_lock import ContextLock +from miles.utils.ft_utils.health_checker import ( + ActiveAndEpoch, + ActivenessTracker, + NoopHealthChecker, + SimpleHealthChecker, + SimpleHealthCheckerConfig, +) +from miles.utils.test_utils.clock import FakeClock + +pytestmark = pytest.mark.usefixtures("dispose_tracked_server_cells") + +_ENDPOINT_CALLS: list[tuple[str, str]] = [] + + +def _make_meta(*, needs_offload: bool = False) -> ServerCellMetadata: + return ServerCellMetadata( + model_id="default", + worker_type="regular", + cell_id="inference-engine-0-0-0", + num_gpus_per_engine=1, + gpu_offset=0, + sglang_api_key=None, + worker_name="inference-engine-0-0-0-0", + needs_offload=needs_offload, + update_weights=True, + workers_hash="pseudo-hash-0", + ) + + +def _make_cell(*, ft_components: list[str], global_activeness: bool = True) -> ServerCell: + return track_server_cell( + ServerCell( + args=make_args(ft_components=ft_components), + meta=_make_meta(), + router_api_client=MagicMock(), + global_health_checker_activeness=lambda: ActiveAndEpoch(active=global_activeness, epoch=0), + ) + ) + + +def _addr_info() -> CellAddrInfo: + return CellAddrInfo(server_url="http://10.0.0.1:30000", bootstrap_port=None, gate_url="http://10.0.0.1:31000") + + +class _RecordingApiClient: + def __init__(self, server_url: str) -> None: + self.server_url = server_url + + async def health_generate(self, timeout: float = 5.0) -> bool: + _ENDPOINT_CALLS.append(("health_generate", self.server_url)) + return True + + +class _NoopRouterApiClient: + async def add_worker(self, **kwargs: Any) -> None: + pass + + async def remove_worker(self, **kwargs: Any) -> None: + pass + + +class TestRolloutCellHealthCheckerGating: + async def test_a_cell_gets_no_checker_when_rollout_ft_is_off(self): + """Probing engines nobody will heal only produces noise and load.""" + cell = _make_cell(ft_components=["train"]) + assert isinstance(cell._health_checker, NoopHealthChecker) + + async def test_a_cell_gets_a_real_checker_when_rollout_ft_is_on(self): + """Rollout healing needs liveness, so the checker must actually be wired up.""" + cell = _make_cell(ft_components=["rollout"]) + assert isinstance(cell._health_checker, SimpleHealthChecker) + + async def test_the_checker_never_waits_a_grace_period(self): + """Activeness flips every weight update window, so a grace period would restart forever.""" + cell = _make_cell(ft_components=["rollout"]) + assert cell._health_checker._config.first_wait == 0.0 + + +class TestRolloutCellHealthCheckerActiveness: + @pytest.mark.parametrize( + "state, expected", + [ + (StateUninitialized(), False), + (StatePendingWeights(addr_info=_addr_info()), True), + (StateServing(addr_info=_addr_info()), True), + ], + ) + async def test_only_a_started_engine_is_probed(self, state, expected): + """An engine whose process is not up yet would fail every probe and look unhealthy.""" + cell = _make_cell(ft_components=["rollout"]) + cell._state = state + assert cell._health_checker._get_activeness().active is expected + + async def test_the_global_flag_can_silence_a_serving_cell(self): + """During a weight update the engine is offloaded, so probing it would kill a healthy cell.""" + cell = _make_cell(ft_components=["rollout"], global_activeness=False) + cell._state = StateServing(addr_info=_addr_info()) + assert cell._health_checker._get_activeness().active is False + + +class TestRolloutCellHealthCheckerProbeEndpoint: + async def test_the_probe_hits_health_generate_on_the_cells_own_engine(self, monkeypatch): + """A probe aimed at another endpoint or another engine would never notice this engine dying.""" + _ENDPOINT_CALLS.clear() + monkeypatch.setattr(server_cell_module, "SGLangApiClient", _RecordingApiClient) + cell = _make_cell(ft_components=["rollout"]) + cell._state = StateServing(addr_info=_addr_info()) + + await cell._health_checker._check_fn() + + assert _ENDPOINT_CALLS == [("health_generate", "http://10.0.0.1:30000")] + + +class TestRolloutCellHealthCheckerPauseBarrier: + async def test_pausing_the_controller_makes_the_predicate_false_for_every_cell(self): + """The whole point of the pause is that no cell is probed inside the protected window.""" + controller, cell = await _make_controller_with_serving_cell() + assert cell._health_checker._get_activeness().active is True + + await controller.offload() + + assert cell._health_checker._get_activeness().active is False + + async def test_pausing_discards_a_probe_that_was_already_in_flight(self): + """A probe launched before the pause must not publish a failure inside the protected window.""" + results: list[bool] = [] + probe_started = asyncio.Event() + + async def _hanging_check() -> None: + probe_started.set() + await asyncio.sleep(3600) + + controller, cell = await _make_controller_with_serving_cell() + checker = _restart_checker_with(cell, check_fn=_hanging_check, on_result=results.append) + await asyncio.wait_for(probe_started.wait(), timeout=1) + + await controller.offload() + + assert checker._probe_task is None + assert results == [] + + async def test_the_pause_returns_only_after_the_in_flight_probe_is_really_gone(self): + """Returning while the probe is still running is exactly the race the barrier has to close.""" + probe_started = asyncio.Event() + probe_cancelled = asyncio.Event() + + async def _hanging_check() -> None: + probe_started.set() + try: + await asyncio.sleep(3600) + except asyncio.CancelledError: + probe_cancelled.set() + raise + + controller, cell = await _make_controller_with_serving_cell() + _restart_checker_with(cell, check_fn=_hanging_check, on_result=None) + await asyncio.wait_for(probe_started.wait(), timeout=1) + + await controller.offload() + + assert probe_cancelled.is_set() + + +class TestRolloutCellActiveAndEpoch: + async def test_reading_the_active_and_epoch_twice_returns_the_very_same_value(self): + """A pull that samples into a tracker of its own degrades the epoch back into a boolean read.""" + _, cell = await _make_controller_with_serving_cell() + + first = cell._get_health_checker_active_and_epoch() + second = cell._get_health_checker_active_and_epoch() + + assert first == second == ActiveAndEpoch(active=True, epoch=0) + + async def test_a_pause_resume_window_completed_between_two_polls_resets_the_failure_counter(self): + """The controller opens and closes the whole window while the loop sleeps, so only the epoch reveals it.""" + controller, cell = await _make_controller_with_serving_cell() + checker, clock = _make_fake_clock_checker(cell) + + checker.start() + await _settle(clock) + await clock.elapse(100.0) + await clock.elapse(10.0) + assert checker._consecutive_failures == 2 + + await controller.offload() + await controller.prepare_eval() + await clock.elapse(10.0) + + assert checker._consecutive_failures == 0 + checker.stop() + + +class TestRolloutCellHealthCheckerDisposal: + async def test_disposing_a_cell_stops_its_checker(self): + """A removed cell whose loop keeps polling leaks the task and the whole cell it closes over.""" + cell = _make_cell(ft_components=["rollout"]) + assert cell._health_checker._task is not None + + await cell.dispose() + + assert cell._health_checker._task is None + + async def test_a_cell_dropped_without_dispose_complains_on_collection(self): + """Silently dropping a cell leaks its checker task forever, so the mistake must be loud.""" + cell = _make_cell(ft_components=["rollout"]) + + with pytest.raises(AssertionError, match="without dispose"): + cell.__del__() + + async def test_a_disposed_cell_is_collected_quietly(self): + """The collection guard must not fire on the normal teardown path.""" + cell = _make_cell(ft_components=["rollout"]) + await cell.dispose() + + cell.__del__() + + +async def _make_controller_with_serving_cell() -> tuple[InferenceController, ServerCell]: + args: Any = make_args(ft_components=["rollout"], colocate=True) + controller = InferenceController.__new__(InferenceController) + controller.args = args + controller.context_lock = ContextLock("InferenceController") + controller._health_checker_activeness = ActivenessTracker(active=True) + + srv = RolloutServer( + server_cells={}, + args=args, + context_lock=controller.context_lock, + global_health_checker_activeness=controller._health_checker_activeness.get, + ) + controller.servers = {"default": srv} + + async with controller.context_lock: + await srv.add_cell(_make_meta(needs_offload=True)) + + cell: ServerCell = track_server_cell(srv.server_cells["inference-engine-0-0-0"]) + cell.router_api_client = _NoopRouterApiClient() + cell._state = StateServing(addr_info=_addr_info()) + return controller, cell + + +def _make_fake_clock_checker(cell: ServerCell) -> tuple[SimpleHealthChecker, FakeClock]: + cell._health_checker.stop() + + async def _failing_check() -> None: + raise RuntimeError("engine down") + + clock = FakeClock() + checker = SimpleHealthChecker( + name=f"rollout-cell-{cell.meta.cell_id}", + check_fn=_failing_check, + get_activeness=cell._get_health_checker_active_and_epoch, + config=SimpleHealthCheckerConfig(interval=10.0, timeout=5.0, first_wait=100.0, failure_threshold=3), + clock=clock, + ) + return checker, clock + + +async def _settle(clock: FakeClock) -> None: + for _ in range(1000): + if clock.pending_count >= 1: + return + await asyncio.sleep(0) + + +def _restart_checker_with(cell: ServerCell, *, check_fn: Any, on_result: Any) -> SimpleHealthChecker: + checker: SimpleHealthChecker = cell._health_checker + checker.stop() + checker._check_fn = check_fn + checker._on_result = on_result + checker.start() + return checker diff --git a/tests/fast/ray/rollout/test_server_cell_state_machine.py b/tests/fast/ray/rollout/test_server_cell_state_machine.py index 2a4be58049b..e5d333c7202 100644 --- a/tests/fast/ray/rollout/test_server_cell_state_machine.py +++ b/tests/fast/ray/rollout/test_server_cell_state_machine.py @@ -1,7 +1,7 @@ from __future__ import annotations import pytest -from tests.fast.ray.rollout.conftest import make_args +from tests.fast.ray.rollout.conftest import make_args, track_server_cell from miles.ray.rollout import server_cell as server_cell_module from miles.ray.rollout.cell_state import ( @@ -14,6 +14,8 @@ ) from miles.ray.rollout.server_cell import ServerCell, ServerCellMetadata +pytestmark = pytest.mark.usefixtures("dispose_tracked_server_cells") + _ADDR_INFO = CellAddrInfo( server_url="http://10.0.0.1:30000", bootstrap_port=None, @@ -96,10 +98,12 @@ async def _probe(server_url: str, api_key, timeout: float = 5.0) -> bool: def _make_cell( *, router: _RecordingRouterApiClient | None = None, args_overrides=None, **meta_overrides ) -> ServerCell: - return ServerCell( - args=make_args(**(args_overrides or {})), - meta=_make_meta(**meta_overrides), - router_api_client=router or _RecordingRouterApiClient(), + return track_server_cell( + ServerCell( + args=make_args(**(args_overrides or {})), + meta=_make_meta(**meta_overrides), + router_api_client=router or _RecordingRouterApiClient(), + ) ) diff --git a/tests/fast/ray/train/test_cell_monitor.py b/tests/fast/ray/train/test_cell_monitor.py index 068e91fd677..f51888541c7 100644 --- a/tests/fast/ray/train/test_cell_monitor.py +++ b/tests/fast/ray/train/test_cell_monitor.py @@ -12,7 +12,7 @@ StateStopped, ) from miles.utils.ft_utils.api_server.models import TriState -from miles.utils.ft_utils.health_checker import ActivenessState, SimpleHealthCheckerConfig +from miles.utils.ft_utils.health_checker import ActiveAndEpoch, SimpleHealthCheckerConfig from miles.utils.ft_utils.indep_dp import IndepDPInfo @@ -125,7 +125,7 @@ async def test_rpc_returns_means_healthy_regardless_of_progress(self): checker = create_trainer_cell_health_checker( cell=cell, config=SimpleHealthCheckerConfig(interval=10.0, timeout=10.0, first_wait=0.0, failure_threshold=3), - get_activeness=lambda: ActivenessState(active=True, epoch=0), + get_activeness=lambda: ActiveAndEpoch(active=True, epoch=0), ) await checker._check_fn() @@ -141,7 +141,7 @@ async def test_rpc_error_propagates_as_unhealthy(self): checker = create_trainer_cell_health_checker( cell=cell, config=SimpleHealthCheckerConfig(interval=10.0, timeout=10.0, first_wait=0.0, failure_threshold=3), - get_activeness=lambda: ActivenessState(active=True, epoch=0), + get_activeness=lambda: ActiveAndEpoch(active=True, epoch=0), ) with pytest.raises(ray.exceptions.RayActorError): @@ -156,7 +156,7 @@ async def test_not_alive_cell_skips_rpc(self): checker = create_trainer_cell_health_checker( cell=cell, config=SimpleHealthCheckerConfig(interval=10.0, timeout=10.0, first_wait=0.0, failure_threshold=3), - get_activeness=lambda: ActivenessState(active=True, epoch=0), + get_activeness=lambda: ActiveAndEpoch(active=True, epoch=0), ) await checker._check_fn() diff --git a/tests/fast/utils/test_arguments.py b/tests/fast/utils/test_arguments.py index 2093c78c2e0..0654cd83b15 100644 --- a/tests/fast/utils/test_arguments.py +++ b/tests/fast/utils/test_arguments.py @@ -1241,6 +1241,10 @@ def test_the_resolved_rollout_config_matches_the_parsed_arguments(self): assert (config.interval, config.timeout, config.first_wait) == (30.0, 30.0, 600.0) + def test_a_rollout_cell_is_reported_unhealthy_on_the_very_first_failed_probe(self): + """At a 30s interval the shared three-failure debounce would hide a dead engine for 90s.""" + assert self._parse([]).rollout_health_check_failure_threshold == 1 + def test_the_trainer_heartbeat_keeps_its_own_debounce(self): """The rollout default must not be pushed down into the shared config: a trainer heartbeat shares an RPC channel with the train step, so one slow reply is a blip, not a dead cell.""" diff --git a/tests/fast/utils/test_health_checker.py b/tests/fast/utils/test_health_checker.py index c6252a41180..1511df94d14 100644 --- a/tests/fast/utils/test_health_checker.py +++ b/tests/fast/utils/test_health_checker.py @@ -1,11 +1,17 @@ +import argparse import asyncio +import pytest +from _pytest.recwarn import WarningsRecorder + +from miles.utils.arguments import get_miles_extra_args_provider from miles.utils.ft_utils.api_server.models import TriState from miles.utils.ft_utils.health_checker import ( - ActivenessState, + ActiveAndEpoch, ActivenessTracker, NoopHealthChecker, SimpleHealthChecker, + SimpleHealthCheckerConfig, ) from miles.utils.test_utils.clock import FakeClock @@ -62,7 +68,7 @@ def active(self) -> bool: def active(self, value: bool) -> None: self._tracker.bump_active(value) - def __call__(self) -> ActivenessState: + def __call__(self) -> ActiveAndEpoch: return self._tracker.get() @@ -657,7 +663,157 @@ async def check_fn() -> None: checker.stop() +class TestCancelInflightProbe: + def _hanging_check_fn(self, started: asyncio.Event): + async def check_fn() -> None: + started.set() + await asyncio.sleep(3600) + + return check_fn + + async def test_a_cancelled_probe_publishes_no_result(self): + """A probe that outlives its window would report a failure about an engine nobody was watching.""" + results: list[bool] = [] + started = asyncio.Event() + checker, _ = _make_checker( + check_fn=self._hanging_check_fn(started), on_result=lambda s: results.append(s), interval=5.0 + ) + checker.start() + await asyncio.wait_for(started.wait(), timeout=1) + + await checker.cancel_inflight_probe() + + assert results == [] + assert checker.status == TriState.UNKNOWN + checker.stop() + + async def test_it_returns_only_after_the_probe_task_is_gone(self): + """Returning while the probe is still running is exactly the race a barrier has to close.""" + started = asyncio.Event() + checker, _ = _make_checker(check_fn=self._hanging_check_fn(started), interval=5.0) + checker.start() + await asyncio.wait_for(started.wait(), timeout=1) + probe_task = checker._probe_task + + await checker.cancel_inflight_probe() + + assert probe_task.cancelled() + checker.stop() + + async def test_the_loop_keeps_polling_after_its_probe_was_cancelled(self): + """Cancelling one probe must not silently kill the checker for the rest of the run.""" + call_count = 0 + started = asyncio.Event() + + async def check_fn() -> None: + nonlocal call_count + call_count += 1 + if call_count == 1: + started.set() + await asyncio.sleep(3600) + + checker, clock = _make_checker(check_fn=check_fn, interval=5.0) + checker.start() + await asyncio.wait_for(started.wait(), timeout=1) + await checker.cancel_inflight_probe() + await _settle(clock) + + await clock.elapse(5.0) + + assert call_count == 2 + checker.stop() + + async def test_cancelling_while_nothing_is_in_flight_is_a_noop(self): + """The controller pauses whether or not a probe happens to be running right then.""" + checker, _ = _make_checker(interval=5.0) + + await checker.cancel_inflight_probe() + + assert checker._probe_task is None + + async def test_stopping_also_kills_a_probe_still_in_flight(self): + """A probe left running after stop() outlives the cell and keeps dialing a dead engine.""" + started = asyncio.Event() + checker, _ = _make_checker(check_fn=self._hanging_check_fn(started), interval=5.0) + checker.start() + await asyncio.wait_for(started.wait(), timeout=1) + probe_task = checker._probe_task + + checker.stop() + + with pytest.raises(asyncio.CancelledError): + await probe_task + assert probe_task.cancelled() + + async def test_stopping_before_a_probe_starts_does_not_abandon_its_coroutine( + self, recwarn: WarningsRecorder + ) -> None: + """Stopping a newly scheduled probe must not leave its coroutine unawaited.""" + checker, _ = _make_checker(interval=5.0) + run_probe = asyncio.create_task(checker._run_probe()) + await asyncio.sleep(0) + + checker.stop() + + with pytest.raises(asyncio.CancelledError): + await run_probe + assert not [warning for warning in recwarn if "was never awaited" in str(warning.message)] + + class TestNoopHealthChecker: def test_noop_status_is_always_unknown(self): checker = NoopHealthChecker() assert checker.status == TriState.UNKNOWN + + +class TestFailureThresholdArgument: + def _parse(self, extra: list[str], **add_arguments_kwargs: int) -> argparse.Namespace: + parser = argparse.ArgumentParser() + SimpleHealthCheckerConfig.add_arguments(parser, prefix="demo-check", **add_arguments_kwargs) + return parser.parse_args(extra) + + def test_a_caller_can_ask_for_a_debounce_of_its_own(self): + """A checker polling on a long interval needs to report the first failure, not the third.""" + args = self._parse([], failure_threshold_default=1) + + assert args.demo_check_failure_threshold == 1 + + def test_a_caller_that_asks_for_nothing_keeps_the_shared_debounce(self): + """Giving one caller a tighter default must not tighten it for every other checker.""" + args = self._parse([]) + + assert args.demo_check_failure_threshold == 3 + + def test_an_explicit_flag_still_beats_a_caller_supplied_default(self): + """A caller default that is nailed into the parser would make the flag unusable.""" + args = self._parse(["--demo-check-failure-threshold", "5"], failure_threshold_default=1) + + assert args.demo_check_failure_threshold == 5 + + +class TestShippedRolloutConfig: + async def test_a_dead_engine_is_reported_unhealthy_after_a_single_probe_under_the_shipped_defaults(self): + """A cell whose engine is already gone must not read Healthy for two more 30s probe intervals.""" + parser = argparse.ArgumentParser() + get_miles_extra_args_provider()(parser) + config = SimpleHealthCheckerConfig.from_args( + parser.parse_args(["--rollout-batch-size", "64"]), prefix="rollout_health_check" + ) + + async def check_fn() -> None: + raise RuntimeError("engine down") + + clock = FakeClock() + checker = SimpleHealthChecker( + name="rollout-cell", + check_fn=check_fn, + get_activeness=_Activeness(), + config=config, + clock=clock, + ) + checker.start() + await _settle(clock) + + assert checker.status == TriState.FALSE + assert checker._consecutive_failures == 1 + checker.stop()