Gate Ray fault injection between weight updates (#3066)

This commit is contained in:
fzyzcjy
2026-09-26 20:41:21 +08:00
committed by GitHub
parent 36ce06f8a3
commit 2098042990
5 changed files with 316 additions and 263 deletions
+16 -13
View File
@@ -29,7 +29,6 @@ from miles.utils.init_once import InitOnce, init_once
from miles.utils.logging_utils import configure_logger
from miles.utils.misc import SimpleTicker
from miles.utils.test_utils.fault_injector import FailureMode
from miles.utils.workers.ray_worker_manager import RayWorkerManager
from miles.utils.workers.registration.hub import RegistrationHub
from miles.utils.workers.registration.models import RegistrationSnapshot
from miles.utils.workers.worker_provider.base import BaseWorkerProvider, CellInfo, StopWatchFn
@@ -94,6 +93,22 @@ class InferenceController:
async def stop_cell_between_weight_updates(self, cell_id: str) -> None:
await self._engine_provider.stop_cells(cell_ids=[cell_id])
# TEMPORARY: exists only so fault injection can take this lock, reverted with the weight-update fault tolerance work
@with_lock
async def inject_fault_between_weight_updates(self, cell_id: str, *, mode: FailureMode, sub_index: int) -> None:
# TEMPORARY: colocate cannot kill rollout workers while trainer ranks own the shared GPUs
server = next((srv for srv in self.servers.values() if cell_id in srv.server_cells), None)
if server is None:
raise KeyError(f"Unknown rollout cell {cell_id!r}")
if not server.health_checker_activeness.get().active:
raise RuntimeError(f"Rollout cell {cell_id!r} is offloaded; refusing fault injection")
await self._engine_provider._worker_manager_handle.inject_fault.remote(
cell_id,
mode=mode.value,
worker_in_cell_index=sub_index,
)
# -------------------------- take over -----------------------------
@lock_exempt
@@ -285,18 +300,6 @@ class InferenceController:
f"Pass model_id to update exactly one of them."
)
# -------------------------- cell operations -----------------------------
# TEMPORARY: exists only so fault injection can take this lock, reverted with the weight-update fault tolerance work
@with_lock
async def inject_fault_between_weight_updates(self, cell_id: str, *, mode: FailureMode, sub_index: int) -> None:
# TEMPORARY: colocate cannot kill rollout workers while trainer ranks own the shared GPUs
if not self._health_checker_activeness.get().active:
raise RuntimeError(f"Rollout cell {cell_id!r} is offloaded; refusing fault injection")
await RayWorkerManager.get_handle().inject_fault.remote(
cell_id, mode=mode.value, worker_in_cell_index=sub_index
)
# -------------------------- eval fleet -----------------------------
@lock_exempt
+27 -5
View File
@@ -27,14 +27,36 @@ class RayCellOperations(BaseCellOperations):
return await self._worker_manager_handle.get_cell_infos.remote(pool_ids=pool_ids)
async def suspend(self, *, cell_id: str) -> None:
if _is_trainer_cell_id(cell_id):
await self._worker_manager_handle.stop_cells.remote([cell_id])
return
# TEMPORARY: taking the lock the weight update holds, reverted with that fault tolerance work
# await self._worker_manager_handle.stop_cells.remote([cell_id]) # use this later
if self._inference_controller is None:
self._inference_controller = self._resolve_inference_controller()
await self._inference_controller.stop_cell_between_weight_updates(cell_id=cell_id)
await self._controller().stop_cell_between_weight_updates(cell_id=cell_id)
async def resume(self, *, cell_id: str) -> None:
await self._worker_manager_handle.start_cells.remote([cell_id])
async def inject_fault(self, *, cell_id: str, mode: FailureMode, sub_index: int) -> None:
await self._worker_manager_handle.inject_fault.remote(cell_id, mode=mode.value, worker_in_cell_index=sub_index)
if _is_trainer_cell_id(cell_id):
await self._worker_manager_handle.inject_fault.remote(
cell_id, mode=mode.value, worker_in_cell_index=sub_index
)
return
# TEMPORARY: taking the lock the weight update holds, reverted with that fault tolerance work
await self._controller().inject_fault_between_weight_updates(
cell_id=cell_id,
mode=mode,
sub_index=sub_index,
)
def _controller(self) -> BaseWorkerHandle:
if self._inference_controller is None:
self._inference_controller = self._resolve_inference_controller()
return self._inference_controller
# TEMPORARY: matched by pool prefix only, until the fault tolerance rework routes trainer cells properly
def _is_trainer_cell_id(cell_id: str) -> bool:
return cell_id.startswith("trainer-engine-")
+12 -5
View File
@@ -307,7 +307,7 @@ Healing witness: one heal per target phase, at P+2 (healed = last cell, ckpt src
```
Type: comparison; both sides run the identical command, only the target is wrapped in the
fault injector, through the pipeline's target_side_context hook
Entry: test_rollout_deterministic__kill_rollout__dp2_cp2__colocate.py, ft-long
Entry: test_rollout_deterministic__kill_rollout__dp4__colocate.py, ft-long
Steps: 8 rollouts (NUM_ROLLOUTS)
Requires: mode.has_real_rollout, and ft_components == ("rollout",) exactly
Compare: dumps rel <= 0 (bitwise); metrics rtol=0 / atol=0 over train/* and rollout/*,
@@ -318,12 +318,14 @@ Regime (both sides):
--sglang-attention-backend flashinfer and --deterministic-mode
- --debug-deterministic-collective and scenario_trainer_deterministic's deterministic env vars
- --sglang-disable-radix-cache
- --rollout-health-check-interval 5
- --rollout-health-check-interval 1
Injection (target side only):
1. Rollout cells, seed 42, exponential mean CRASH_INTERVAL_SECONDS (120s)
1. Rollout cells, seed 42, exponential mean CRASH_INTERVAL_SECONDS (30s)
2. Forms drawn per (cluster backend, cell type), as in the soaks
3. Stop the injector with a 5s timeout, then re-use the soak's rollout witnesses: >= 2
3. Stop accepting faults after six completed rollouts, leaving the final two rollouts for recovery
4. Stop the injector, waiting out a mid-flight injection for at most
STOP_AND_JOIN_TIMEOUT_SECONDS (180s), then re-use the soak's rollout witnesses: >= 2
accepted rollout injections, each paired with one completed recovery cycle
Assertions:
@@ -337,7 +339,12 @@ Assertions:
- **Why it exists**: an engine dying and being replaced mid-generation is supposed to be invisible to training, and "invisible" is a claim about bits; the rollout soak only ever asserted survival.
- **Why the shared deterministic recipe**: the assertion is deterministic replay across fresh inference engines, not true-on-policy training. Reusing the same FlashInfer recipe as the main deterministic trainer-FT test avoids a second, incompatible attention-backend contract.
- **Why `--sglang-disable-radix-cache`**: a replacement engine serves with a cold prefix cache where the baseline's was warm, and deterministic inference is nowhere documented as prefix-cache-length invariant.
- **Why `--rollout-health-check-interval 5`**: the generation retry loop gives up after ~60s while the default health check needs 90-120s to evict a dead worker, so a request could exhaust its retries against a corpse.
- **Why this recipe disables batch-variant MM fallback**: a rollout worker loss changes co-batching while the pool is healing; permitting an `einsum` fallback would make the same seeded request depend on that temporary batch shape. The scenario injects the environment override without changing the production default.
- **Why `--rollout-health-check-interval 1`**: healthy generation can finish between two five-second polls; the short scenario needs at least one fresh Serving observation before its lock-protected injection attempt.
- **Why this scenario polls the fault window every 0.2 seconds**: colocated generation windows are only a few seconds long, so the generic two-second scheduler cadence can miss every Serving observation in an eight-rollout run.
- **Why one quiescent poll on Ray only**: colocate exposes Serving only between train phases, and on Ray the injection endpoint takes the inference-controller lock and atomically rejects inactive servers, so the generic stable-serving gate is redundant. The Kubernetes forms (`exec_sigkill`, `delete_pod`) act on the pod without that lock, so a stale Serving poll or a slow `kubectl` would land the fault after the colocated engines are already offloaded; Kubernetes therefore keeps the generic 60-poll gate.
- **Why the final two rollouts accept no new fault**: the scheduler keeps observing recovery but closes admission after rollout 5, so teardown cannot race a newly accepted replacement.
- **Why Ray checks Serving again inside the injection lock**: the lock excludes weight-update and offload transitions, while the Serving check also rejects the subsequent colocated trainer phase after `offload()` has released the lock.
- **Why every namespace, not just `train/`**: an engine crash shows up first in `rollout/raw_reward` or `rollout/log_probs`. `perf/` is left out by name, being wall-clock and throughput that a relaunch moves by definition, and a metric in neither namespace fails the run rather than being dropped quietly.
- **Why the weights-moved gate**: bitwise equality is also satisfied by two runs that trained on nothing.
- **Why not a loss or reward curve**: neither is a progress signal here — the reward is `deterministic_random`, a hash of the response, and GRPO's surrogate loss is not monotone even while a run learns. Over eight rollouts neither moves for a reason worth asserting, and the weights either changed or they did not.
@@ -133,29 +133,12 @@ class _RecordingServer:
self.cells_timeouts.append(timeout)
class _AbortRecordingServer:
def __init__(
self,
name: str,
*,
started: list[str],
completed: list[str],
error: Exception | None = None,
scheduling_turns: int = 1,
) -> None:
self.name = name
self.started = started
self.completed = completed
self.error = error
self.scheduling_turns = scheduling_turns
class _RecordingRemoteMethod:
def __init__(self) -> None:
self.calls: list[tuple[tuple[Any, ...], dict[str, Any]]] = []
async def abort_all(self) -> None:
self.started.append(self.name)
for _ in range(self.scheduling_turns):
await asyncio.sleep(0)
if self.error is not None:
raise self.error
self.completed.append(self.name)
async def remote(self, *args: Any, **kwargs: Any) -> None:
self.calls.append((args, kwargs))
class _FakeUpdatableCell:
@@ -191,6 +174,7 @@ class _FakeWorkerProvider(BaseWorkerProvider):
def __init__(self, cell_infos: list[CellInfo], *, pool_ids: list[str] | None = None) -> None:
self._cell_infos = cell_infos
self._pools = pool_ids or []
self._worker_manager_handle = SimpleNamespace(inject_fault=_RecordingRemoteMethod())
self.watched_pool_ids: list[str] | None = None
self.initialized = False
self.stop_watch_calls = 0
@@ -252,37 +236,6 @@ def _make_controller(
return controller
class TestAbortAll:
async def test_abort_all_reaches_every_server_before_propagating_the_first_failure(self) -> None:
"""Every server is aborted to completion before the first fleet failure is propagated."""
started: list[str] = []
completed: list[str] = []
first_failure = RuntimeError("first server refused abort")
controller = _make_controller(
{
"first": _AbortRecordingServer(
"first",
started=started,
completed=completed,
error=first_failure,
),
"second": _AbortRecordingServer(
"second",
started=started,
completed=completed,
scheduling_turns=2,
),
}
)
with pytest.raises(RuntimeError) as exc_info:
await controller.abort_all()
assert exc_info.value is first_failure
assert set(started) == {"first", "second"}
assert completed == ["second"]
class TestHealthCheckerActiveness:
@pytest.mark.asyncio
async def test_offload_pauses_probing_before_putting_engines_to_sleep(self):
@@ -346,6 +299,63 @@ class TestHealthCheckerActiveness:
assert srv.health_checker_activeness.get().active
class TestRolloutFaultInjectionWindow:
async def test_fault_injection_refuses_an_unknown_rollout_cell(self) -> None:
"""An unknown rollout cell raises its identifier without reaching the worker manager."""
cell_id = "unknown-rollout-cell"
provider = _FakeWorkerProvider([])
controller = _make_controller(
{"default": _RecordingServer(server_cells={"inference-engine-0-0-0": object()})},
engine_provider=provider,
)
with pytest.raises(KeyError) as exc_info:
await controller.inject_fault_between_weight_updates(
cell_id=cell_id,
mode=FailureMode.SIGKILL,
sub_index=0,
)
assert cell_id in str(exc_info.value)
assert provider._worker_manager_handle.inject_fault.calls == []
@pytest.mark.asyncio
async def test_fault_injection_reaches_a_serving_rollout_cell(self) -> None:
"""A serving rollout cell accepts a fault through the Ray worker manager."""
cell_id = "inference-engine-0-0-0"
server = _RecordingServer(server_cells={cell_id: object()})
provider = _FakeWorkerProvider([])
controller = _make_controller({"default": server}, engine_provider=provider)
await controller.inject_fault_between_weight_updates(
cell_id=cell_id,
mode=FailureMode.SIGKILL,
sub_index=0,
)
assert provider._worker_manager_handle.inject_fault.calls == [
((cell_id,), {"mode": "sigkill", "worker_in_cell_index": 0})
]
@pytest.mark.asyncio
async def test_fault_injection_refuses_an_offloaded_rollout_cell(self) -> None:
"""Colocate must not kill rollout processes while trainer ranks own the shared GPUs."""
cell_id = "inference-engine-0-0-0"
server = _RecordingServer(server_cells={cell_id: object()})
server.health_checker_activeness.bump_active(False)
provider = _FakeWorkerProvider([])
controller = _make_controller({"default": server}, engine_provider=provider)
with pytest.raises(RuntimeError, match="is offloaded; refusing fault injection"):
await controller.inject_fault_between_weight_updates(
cell_id=cell_id,
mode=FailureMode.SIGKILL,
sub_index=0,
)
assert provider._worker_manager_handle.inject_fault.calls == []
class TestReconcile:
@pytest.fixture
def servers(self) -> dict[str, _RecordingServer]:
@@ -642,9 +652,12 @@ class TestInitSubscription:
assert "session-server" not in provider.watched_pool_ids
@pytest.mark.asyncio
async def test_init_survives_a_router_cell_offered_by_the_provider(self, monkeypatch: pytest.MonkeyPatch):
"""A router cell carries no engine meta, so a too-wide subscription kills startup in the initial sync."""
async def test_init_subscribes_narrowly_enough_to_never_see_a_router_cell(self, monkeypatch: pytest.MonkeyPatch):
"""A router cell carries no engine meta, so reading one as engine meta would kill startup; the
controller is safe because it subscribes to the engine pools alone."""
args = make_args()
assert compute_router_pool_id(0) not in compute_engine_pool_ids(args)
router_info = CellInfo(
cell_id="inference-router-0-0",
pool_id=compute_router_pool_id(0),
@@ -1158,194 +1171,124 @@ async def _raise_async(cell: ServerCell) -> None:
raise RuntimeError("injected init failure")
class _RecordingWorkerManager:
def __init__(self) -> None:
_ROLLOUT_CELL_ID = "inference-engine-0-0-0"
class _StoppingWorkerProvider(_FakeWorkerProvider):
def __init__(self, *, completion: asyncio.Future | None = None) -> None:
super().__init__([])
self.stopped_cells: list[list[str]] = []
self.injected: list[tuple[str, dict[str, Any]]] = []
self.stop_requested = asyncio.Event()
self._completion = completion
@property
def stop_cells(self) -> Any:
return _RecordingRemoteCall(lambda cell_ids: self.stopped_cells.append(list(cell_ids)))
@property
def inject_fault(self) -> Any:
return _RecordingRemoteCall(lambda cell_id, **kwargs: self.injected.append((cell_id, kwargs)))
async def stop_cells(self, *, cell_ids: list[str]) -> None:
self.stopped_cells.append(list(cell_ids))
self.stop_requested.set()
if self._completion is not None:
await self._completion
class _RecordingRemoteCall:
def __init__(self, record: Any) -> None:
self._record = record
def remote(self, *args: Any, **kwargs: Any) -> asyncio.Future:
self._record(*args, **kwargs)
future: asyncio.Future = asyncio.get_event_loop().create_future()
future.set_result(None)
return future
def _make_cell_operations_controller(
provider: _StoppingWorkerProvider, *, probing_paused: bool = False
) -> InferenceController:
server = _RecordingServer(server_cells={_ROLLOUT_CELL_ID: object()})
if probing_paused:
server.health_checker_activeness.bump_active(False)
return _make_controller({"default": server}, engine_provider=provider)
def _patch_worker_manager(monkeypatch: pytest.MonkeyPatch) -> _RecordingWorkerManager:
manager = _RecordingWorkerManager()
monkeypatch.setattr(inference_controller_module, "RayWorkerManager", SimpleNamespace(get_handle=lambda: manager))
return manager
async def _hold_context_lock(controller: InferenceController) -> tuple[asyncio.Task, asyncio.Event]:
entered, may_finish = asyncio.Event(), asyncio.Event()
async def _hold() -> None:
async with controller.context_lock:
entered.set()
await may_finish.wait()
class _BlockingWorkerManager:
def __init__(self, *, completion: asyncio.Future) -> None:
self.requested = asyncio.Event()
self.stopped_cells: list[list[str]] = []
self.completion = completion
@property
def stop_cells(self) -> Any:
return _BlockingRemoteCall(manager=self)
class _BlockingRemoteCall:
def __init__(self, *, manager: _BlockingWorkerManager) -> None:
self._manager = manager
def remote(self, cell_ids: list[str]) -> asyncio.Future:
self._manager.stopped_cells.append(list(cell_ids))
self._manager.requested.set()
return self._manager.completion
def _patch_blocking_worker_manager(
monkeypatch: pytest.MonkeyPatch, *, completion: asyncio.Future
) -> _BlockingWorkerManager:
manager = _BlockingWorkerManager(completion=completion)
monkeypatch.setattr(inference_controller_module, "RayWorkerManager", SimpleNamespace(get_handle=lambda: manager))
return manager
holder = asyncio.create_task(_hold())
await entered.wait()
return holder, may_finish
class TestCellOperations:
@pytest.mark.asyncio
async def test_stop_cell_between_weight_updates_is_forwarded_to_the_worker_manager(
self, monkeypatch: pytest.MonkeyPatch
):
"""The manager owns the processes, so the controller only serializes the suspension."""
manager = _patch_worker_manager(monkeypatch)
controller = _make_controller({})
async def test_stop_cell_between_weight_updates_is_forwarded_to_the_engine_provider(self):
"""The provider owns the processes, so the controller only serializes the suspension."""
provider = _StoppingWorkerProvider()
controller = _make_cell_operations_controller(provider)
await controller.stop_cell_between_weight_updates("engine-0")
await controller.stop_cell_between_weight_updates(_ROLLOUT_CELL_ID)
assert manager.stopped_cells == [["engine-0"]]
assert provider.stopped_cells == [[_ROLLOUT_CELL_ID]]
@pytest.mark.asyncio
async def test_inject_fault_between_weight_updates_is_forwarded_to_the_worker_manager(
self, monkeypatch: pytest.MonkeyPatch
):
"""Injection targets one worker of one cell, and only the manager can reach it."""
manager = _patch_worker_manager(monkeypatch)
controller = _make_controller({})
await controller.inject_fault_between_weight_updates("engine-0", mode=FailureMode.SIGKILL, sub_index=1)
assert manager.injected == [("engine-0", {"mode": "sigkill", "worker_in_cell_index": 1})]
@pytest.mark.asyncio
async def test_inject_fault_between_weight_updates_is_refused_while_probing_is_paused(
self, monkeypatch: pytest.MonkeyPatch
):
"""An offloaded or updating cell would report a crash that the trainer cannot distinguish from its own pause."""
manager = _patch_worker_manager(monkeypatch)
controller = _make_controller({})
controller._health_checker_activeness.bump_active(False)
with pytest.raises(RuntimeError, match="refusing fault injection"):
await controller.inject_fault_between_weight_updates("engine-0", mode=FailureMode.SIGKILL, sub_index=0)
assert manager.injected == []
@pytest.mark.asyncio
async def test_stop_cell_between_weight_updates_waits_until_the_weight_update_window_closes(
self, monkeypatch: pytest.MonkeyPatch
):
async def test_stop_cell_between_weight_updates_waits_until_the_weight_update_window_closes(self):
"""Suspending a cell mid-broadcast leaves the trainer waiting on an engine that is being torn down."""
manager = _patch_worker_manager(monkeypatch)
controller = _make_controller({})
entered = asyncio.Event()
may_finish = asyncio.Event()
provider = _StoppingWorkerProvider()
controller = _make_cell_operations_controller(provider)
holder, may_finish = await _hold_context_lock(controller)
async def _hold_lock() -> None:
async with controller.context_lock:
entered.set()
await may_finish.wait()
holder = asyncio.create_task(_hold_lock())
await entered.wait()
stopping = asyncio.create_task(controller.stop_cell_between_weight_updates("engine-0"))
stopping = asyncio.create_task(controller.stop_cell_between_weight_updates(_ROLLOUT_CELL_ID))
await asyncio.sleep(0)
assert manager.stopped_cells == []
assert provider.stopped_cells == []
may_finish.set()
await holder
await stopping
assert manager.stopped_cells == [["engine-0"]]
assert provider.stopped_cells == [[_ROLLOUT_CELL_ID]]
@pytest.mark.asyncio
async def test_inject_fault_between_weight_updates_waits_until_the_weight_update_window_closes(
self, monkeypatch: pytest.MonkeyPatch
):
async def test_inject_fault_between_weight_updates_waits_until_the_weight_update_window_closes(self):
"""Injection racing a broadcast is the same hazard as suspension, so it takes the same turn."""
manager = _patch_worker_manager(monkeypatch)
controller = _make_controller({})
entered = asyncio.Event()
may_finish = asyncio.Event()
provider = _StoppingWorkerProvider()
controller = _make_cell_operations_controller(provider)
holder, may_finish = await _hold_context_lock(controller)
async def _hold_lock() -> None:
async with controller.context_lock:
entered.set()
await may_finish.wait()
holder = asyncio.create_task(_hold_lock())
await entered.wait()
injecting = asyncio.create_task(
controller.inject_fault_between_weight_updates("engine-0", mode=FailureMode.SIGKILL, sub_index=0)
controller.inject_fault_between_weight_updates(_ROLLOUT_CELL_ID, mode=FailureMode.SIGKILL, sub_index=0)
)
await asyncio.sleep(0)
assert manager.injected == []
assert provider._worker_manager_handle.inject_fault.calls == []
may_finish.set()
await holder
await injecting
assert manager.injected == [("engine-0", {"mode": "sigkill", "worker_in_cell_index": 0})]
assert provider._worker_manager_handle.inject_fault.calls == [
((_ROLLOUT_CELL_ID,), {"mode": "sigkill", "worker_in_cell_index": 0})
]
async def test_stop_cell_between_weight_updates_is_allowed_while_probing_is_paused(
self, monkeypatch: pytest.MonkeyPatch
):
@pytest.mark.asyncio
async def test_stop_cell_between_weight_updates_is_allowed_while_probing_is_paused(self):
"""An offloaded cell is the one a heal loop most needs to suspend, so only injection is refused."""
manager = _patch_worker_manager(monkeypatch)
controller = _make_controller({})
controller._health_checker_activeness.bump_active(False)
provider = _StoppingWorkerProvider()
controller = _make_cell_operations_controller(provider, probing_paused=True)
await controller.stop_cell_between_weight_updates("engine-0")
await controller.stop_cell_between_weight_updates(_ROLLOUT_CELL_ID)
assert manager.stopped_cells == [["engine-0"]]
assert provider.stopped_cells == [[_ROLLOUT_CELL_ID]]
async def test_inject_fault_between_weight_updates_refuses_a_pause_that_began_while_it_waited(
self, monkeypatch: pytest.MonkeyPatch
):
@pytest.mark.asyncio
async def test_inject_fault_between_weight_updates_refuses_a_pause_that_began_while_it_waited(self):
"""Reading the pause before taking the lock would kill a cell the offload has since put to sleep."""
manager = _patch_worker_manager(monkeypatch)
controller = _make_controller({})
entered = asyncio.Event()
may_finish = asyncio.Event()
provider = _StoppingWorkerProvider()
controller = _make_cell_operations_controller(provider)
server = controller.servers["default"]
entered, may_finish = asyncio.Event(), asyncio.Event()
async def _pause_probing_under_the_lock() -> None:
async with controller.context_lock:
entered.set()
await may_finish.wait()
controller._health_checker_activeness.bump_active(False)
server.health_checker_activeness.bump_active(False)
holder = asyncio.create_task(_pause_probing_under_the_lock())
await entered.wait()
injecting = asyncio.create_task(
controller.inject_fault_between_weight_updates("engine-0", mode=FailureMode.SIGKILL, sub_index=0)
controller.inject_fault_between_weight_updates(_ROLLOUT_CELL_ID, mode=FailureMode.SIGKILL, sub_index=0)
)
await asyncio.sleep(0)
may_finish.set()
@@ -1354,38 +1297,39 @@ class TestCellOperations:
with pytest.raises(RuntimeError, match="refusing fault injection"):
await injecting
assert manager.injected == []
assert provider._worker_manager_handle.inject_fault.calls == []
async def test_a_refused_injection_leaves_the_weight_update_lock_free(self, monkeypatch: pytest.MonkeyPatch):
@pytest.mark.asyncio
async def test_a_refused_injection_leaves_the_weight_update_lock_free(self):
"""A refusal that kept the lock would hang the next weight update instead of only skipping the injection."""
manager = _patch_worker_manager(monkeypatch)
controller = _make_controller({})
controller._health_checker_activeness.bump_active(False)
provider = _StoppingWorkerProvider()
controller = _make_cell_operations_controller(provider, probing_paused=True)
with pytest.raises(RuntimeError, match="refusing fault injection"):
await controller.inject_fault_between_weight_updates("engine-0", mode=FailureMode.SIGKILL, sub_index=0)
await controller.inject_fault_between_weight_updates(
_ROLLOUT_CELL_ID, mode=FailureMode.SIGKILL, sub_index=0
)
assert not controller.context_lock.locked
await controller.stop_cell_between_weight_updates("engine-0")
await controller.stop_cell_between_weight_updates(_ROLLOUT_CELL_ID)
assert manager.stopped_cells == [["engine-0"]]
assert provider.stopped_cells == [[_ROLLOUT_CELL_ID]]
async def test_a_weight_update_cannot_start_while_a_suspension_is_still_running(
self, monkeypatch: pytest.MonkeyPatch
):
"""Releasing the lock before the manager has torn the cell down reopens the very race this serializes."""
@pytest.mark.asyncio
async def test_a_weight_update_cannot_start_while_a_suspension_is_still_running(self):
"""Releasing the lock before the provider has torn the cell down reopens the very race this serializes."""
completion: asyncio.Future = asyncio.get_running_loop().create_future()
manager = _patch_blocking_worker_manager(monkeypatch, completion=completion)
controller = _make_controller({})
provider = _StoppingWorkerProvider(completion=completion)
controller = _make_cell_operations_controller(provider)
weight_update_started = asyncio.Event()
async def _start_weight_update() -> None:
async with controller.context_lock:
weight_update_started.set()
stopping = asyncio.create_task(controller.stop_cell_between_weight_updates("engine-0"))
await manager.requested.wait()
stopping = asyncio.create_task(controller.stop_cell_between_weight_updates(_ROLLOUT_CELL_ID))
await provider.stop_requested.wait()
weight_update = asyncio.create_task(_start_weight_update())
await asyncio.sleep(0)
@@ -5,14 +5,20 @@ from dataclasses import dataclass
from types import SimpleNamespace
from typing import Any
import pytest
from miles.ray.rollout.inference_controller import InferenceController
from miles.utils.context_lock import ContextLock
from miles.utils.ft_utils.health_checker import ActivenessTracker
from miles.utils.test_utils.fault_injector import FailureMode
from miles.utils.workers.cell_operations.ray import RayCellOperations
_TRAINER_CELL_ID = "trainer-engine-actor-00001"
class _RecordingEngineProvider:
def __init__(self) -> None:
def __init__(self, *, worker_manager: _RecordingWorkerManagerHandle) -> None:
self._worker_manager_handle = worker_manager
self.stopped: list[str] = []
async def stop_cells(self, *, cell_ids: list[str]) -> None:
@@ -35,6 +41,7 @@ class _RecordingWorkerManagerHandle:
self.calls: list[tuple[str, tuple[Any, ...], dict[str, Any]]] = []
self.get_cell_infos = _RecordingRemoteMethod(name="get_cell_infos", calls=self.calls)
self.start_cells = _RecordingRemoteMethod(name="start_cells", calls=self.calls)
self.stop_cells = _RecordingRemoteMethod(name="stop_cells", calls=self.calls)
self.inject_fault = _RecordingRemoteMethod(name="inject_fault", calls=self.calls)
@@ -47,9 +54,15 @@ class _Fixture:
def _make_fixture() -> _Fixture:
provider = _RecordingEngineProvider()
controller = InferenceController(SimpleNamespace(), engine_provider=provider, router_providers=[])
worker_manager = _RecordingWorkerManagerHandle()
provider = _RecordingEngineProvider(worker_manager=worker_manager)
controller = InferenceController(SimpleNamespace(), engine_provider=provider, router_providers=[])
controller.servers = {
"actor": SimpleNamespace(
server_cells={"engine-0-2": SimpleNamespace()},
health_checker_activeness=ActivenessTracker(active=True),
)
}
return _Fixture(
provider=provider,
controller=controller,
@@ -71,30 +84,6 @@ async def _settle() -> None:
await asyncio.sleep(0)
class TestRayCellOperationsProtocol:
async def test_cell_infos_forwards_pool_ids_and_returns_the_actor_result(self) -> None:
"""Cell info reads forward every pool ID by keyword and preserve the actor result."""
worker_manager = _RecordingWorkerManagerHandle()
operations = RayCellOperations(worker_manager_handle=worker_manager)
pool_ids = ["engine-0", "rollout-1"]
actor_result = {"engine-0-2": object(), "rollout-1-3": object()}
worker_manager.get_cell_infos.result = actor_result
result = await operations.cell_infos(pool_ids=pool_ids)
assert worker_manager.calls == [("get_cell_infos", (), {"pool_ids": pool_ids})]
assert result is actor_result
async def test_resume_starts_exactly_the_requested_cell(self) -> None:
"""Resuming a cell sends exactly that cell ID in a one-element list."""
worker_manager = _RecordingWorkerManagerHandle()
operations = RayCellOperations(worker_manager_handle=worker_manager)
await operations.resume(cell_id="engine-0-2")
assert worker_manager.calls == [("start_cells", (["engine-0-2"],), {})]
async def test_a_suspend_waits_for_the_controller_lock_instead_of_reaching_the_worker_manager() -> None:
"""A suspend arriving mid weight update must not reach the worker manager until the update ends."""
fixture = _make_fixture()
@@ -115,8 +104,78 @@ async def test_a_suspend_waits_for_the_controller_lock_instead_of_reaching_the_w
assert fixture.worker_manager.calls == []
async def test_every_other_operation_goes_straight_through() -> None:
"""Only a suspend takes a rank out of a live collective, so nothing else pays for the lock."""
async def test_inject_fault_waits_for_the_controller_lock() -> None:
"""A fault arriving mid weight update must wait before killing an engine worker."""
fixture = _make_fixture()
acquired, release = asyncio.Event(), asyncio.Event()
holding = asyncio.create_task(_hold_lock(lock=fixture.controller.context_lock, acquired=acquired, release=release))
await acquired.wait()
injecting = asyncio.create_task(
fixture.operations.inject_fault(cell_id="engine-0-2", mode=FailureMode.SIGKILL, sub_index=0)
)
await _settle()
assert not injecting.done()
assert fixture.worker_manager.calls == []
release.set()
await holding
await injecting
assert fixture.worker_manager.calls == [
("inject_fault", ("engine-0-2",), {"mode": "sigkill", "worker_in_cell_index": 0})
]
async def test_a_trainer_cells_fault_reaches_the_worker_manager() -> None:
"""Regression: routing a trainer cell through the rollout controller raised, so the actor never died."""
fixture = _make_fixture()
acquired, release = asyncio.Event(), asyncio.Event()
holding = asyncio.create_task(_hold_lock(lock=fixture.controller.context_lock, acquired=acquired, release=release))
await acquired.wait()
await asyncio.wait_for(
fixture.operations.inject_fault(cell_id=_TRAINER_CELL_ID, mode=FailureMode.SIGKILL, sub_index=0),
timeout=5.0,
)
assert fixture.worker_manager.calls == [
("inject_fault", (_TRAINER_CELL_ID,), {"mode": "sigkill", "worker_in_cell_index": 0})
]
release.set()
await holding
async def test_a_trainer_cells_suspend_reaches_the_worker_manager() -> None:
"""A trainer cell is none of the controller's business, and its lock would only stall the stop."""
fixture = _make_fixture()
acquired, release = asyncio.Event(), asyncio.Event()
holding = asyncio.create_task(_hold_lock(lock=fixture.controller.context_lock, acquired=acquired, release=release))
await acquired.wait()
await asyncio.wait_for(fixture.operations.suspend(cell_id=_TRAINER_CELL_ID), timeout=5.0)
assert fixture.worker_manager.calls == [("stop_cells", ([_TRAINER_CELL_ID],), {})]
assert fixture.provider.stopped == []
release.set()
await holding
async def test_a_rollout_cell_the_controller_does_not_list_yet_still_goes_through_the_controller() -> None:
"""Routing on live membership would kill an engine being replaced without the weight-update lock."""
fixture = _make_fixture()
with pytest.raises(KeyError):
await asyncio.wait_for(
fixture.operations.inject_fault(cell_id="engine-0-7", mode=FailureMode.SIGKILL, sub_index=0), timeout=5.0
)
assert fixture.worker_manager.calls == []
async def test_non_disruptive_operations_go_straight_through() -> None:
"""Cell reads and resumes do not wait for the weight-update lock."""
fixture = _make_fixture()
acquired, release = asyncio.Event(), asyncio.Event()
holding = asyncio.create_task(_hold_lock(lock=fixture.controller.context_lock, acquired=acquired, release=release))
@@ -124,18 +183,36 @@ async def test_every_other_operation_goes_straight_through() -> None:
await asyncio.wait_for(fixture.operations.cell_infos(pool_ids=["engine-0"]), timeout=5.0)
await asyncio.wait_for(fixture.operations.resume(cell_id="engine-0-2"), timeout=5.0)
await asyncio.wait_for(
fixture.operations.inject_fault(cell_id="engine-0-2", mode=FailureMode.SIGKILL, sub_index=0),
timeout=5.0,
)
assert [name for name, _, _ in fixture.worker_manager.calls] == ["get_cell_infos", "start_cells", "inject_fault"]
assert [name for name, _, _ in fixture.worker_manager.calls] == ["get_cell_infos", "start_cells"]
assert fixture.provider.stopped == []
release.set()
await holding
class TestRayCellOperationsProtocol:
async def test_cell_infos_forwards_pool_ids_and_returns_the_actor_result(self) -> None:
"""Cell info reads forward every pool ID by keyword and preserve the actor result."""
fixture = _make_fixture()
pool_ids = ["engine-0", "rollout-1"]
actor_result = {"engine-0-2": SimpleNamespace(), "rollout-1-3": SimpleNamespace()}
fixture.worker_manager.get_cell_infos.result = actor_result
result = await fixture.operations.cell_infos(pool_ids=pool_ids)
assert fixture.worker_manager.calls == [("get_cell_infos", (), {"pool_ids": pool_ids})]
assert result is actor_result
async def test_resume_starts_exactly_the_requested_cell(self) -> None:
"""Resuming a cell sends exactly that cell ID in a one-element list."""
fixture = _make_fixture()
await fixture.operations.resume(cell_id="engine-0-2")
assert fixture.worker_manager.calls == [("start_cells", (["engine-0-2"],), {})]
class TestRayCellOperationsInferenceControllerResolution:
async def test_the_inference_controller_is_resolved_only_when_a_disruptive_operation_needs_it(self) -> None:
"""Construction, reads, and resumes do not resolve the controller before a suspend needs it."""