diff --git a/miles/backends/sglang_utils/sglang_engine.py b/miles/backends/sglang_utils/sglang_engine.py index 60ea2fab1c8..67ef61c7c05 100644 --- a/miles/backends/sglang_utils/sglang_engine.py +++ b/miles/backends/sglang_utils/sglang_engine.py @@ -5,14 +5,11 @@ import os import httpx -import sglang_router -from packaging.version import parse from sglang.srt.server_args import ServerArgs from sglang.srt.utils import kill_process_tree from miles.backends.megatron_utils.lora_utils import convert_target_modules_to_hf, sglang_lora_target_all_sentinel from miles.backends.sglang_utils.sglang_api_client import wait_server_healthy -from miles.backends.sglang_utils.sglang_router_api_client import SGLangRouterApiClient from miles.ray.ray_actor import RayActor from miles.utils import async_utils from miles.utils.env_report import collect_and_print_node_env_report @@ -99,10 +96,6 @@ def build_server_url(host: str, port: int) -> str: return f"http://{format_v6_uri(host)}:{port}" -def _use_legacy_router_api(args) -> bool: - return parse(sglang_router.__version__) <= parse("0.2.1") or args.use_miles_router - - class SGLangEngine(RayActor): def __init__( self, @@ -146,8 +139,6 @@ def init( nccl_port, host=None, disaggregation_bootstrap_port=None, - router_ip=None, - router_port=None, engine_info_bootstrap_port=None, ): if env_report := self.args.env_report: @@ -157,9 +148,6 @@ def init( partial_env_report=env_report, ) - self.router_ip = router_ip if router_ip is not None else self.args.sglang_router_ip - self.router_port = router_port if router_port is not None else self.args.sglang_router_port - host = host or get_host_info()[1] host = format_v6_uri(host) @@ -186,7 +174,6 @@ def init( self.server_port = server_args_dict["port"] self.server_url = build_server_url(self.server_host, self.server_port) - self.router_api_client = SGLangRouterApiClient(router_url=f"http://{self.router_ip}:{self.router_port}") if self.args.rollout_external: self._init_external(server_args_dict, external_engine_need_check_fields=external_engine_need_check_fields) @@ -221,30 +208,11 @@ def _init_normal(self, server_args_dict): server_args = ServerArgs(**{**server_args_dict, "host": server_args_dict["host"].strip("[]")}) self.process = launch_server_process(server_args) - if self.node_rank == 0 and self.router_ip and self.router_port: - async_utils.run( - self.router_api_client.add_worker( - worker_url=self.server_url, - worker_type=self.worker_type, - use_legacy_api=_use_legacy_router_api(self.args), - bootstrap_port=( - server_args_dict["disaggregation_bootstrap_port"] if self.worker_type == "prefill" else None - ), - ) - ) - def shutdown(self): if self.args.rollout_external: return logger.info(f"Shutdown engine {self.server_host}:{self.server_port}...") - if self.node_rank == 0: - async_utils.run( - self.router_api_client.remove_worker( - worker_url=self.server_url, - use_legacy_api=_use_legacy_router_api(self.args), - ) - ) kill_process_tree(self.process.pid) def simulate_crash(self): diff --git a/miles/backends/sglang_utils/sglang_router_api_client.py b/miles/backends/sglang_utils/sglang_router_api_client.py index 5fc1d862351..471d3794425 100644 --- a/miles/backends/sglang_utils/sglang_router_api_client.py +++ b/miles/backends/sglang_utils/sglang_router_api_client.py @@ -10,6 +10,10 @@ logger = logging.getLogger(__name__) +def use_legacy_router_api(args) -> bool: + return parse(sglang_router.__version__) <= parse("0.2.1") or args.use_miles_router + + @dataclasses.dataclass(frozen=True) class SGLangRouterApiClient: router_url: str diff --git a/miles/ray/rollout/rollout_server.py b/miles/ray/rollout/rollout_server.py index e6dc32615cd..104d3f29d51 100644 --- a/miles/ray/rollout/rollout_server.py +++ b/miles/ray/rollout/rollout_server.py @@ -10,6 +10,7 @@ from miles.ray.rollout.router_manager import start_router from miles.ray.rollout.server_engine import ServerEngine from miles.ray.rollout.server_group import ServerGroup +from miles.utils import async_utils logger = logging.getLogger(__name__) @@ -89,6 +90,7 @@ def start_rollout_servers(args, pg) -> dict[str, "RolloutServer"]: for group, new_engine_indices in zip(server_groups, new_engine_indices_per_group, strict=True): group.mark_alive(engine_indices=new_engine_indices) + async_utils.run(group.register_workers(new_engine_indices)) servers[model_cfg.name] = RolloutServer( server_groups=server_groups, diff --git a/miles/ray/rollout/server_engine.py b/miles/ray/rollout/server_engine.py index 089b1a434e9..346a501f66a 100644 --- a/miles/ray/rollout/server_engine.py +++ b/miles/ray/rollout/server_engine.py @@ -27,18 +27,18 @@ def __init__(self): def mark_allocated_uninitialized(self, actor_handle: ray.actor.ActorHandle): self._change_state("mark_allocated", _StateStopped, _StateAllocatedUninitialized(actor_handle=actor_handle)) - def set_server_url(self, server_url: str) -> None: + def set_addressing(self, addr_info: AddrInfo) -> None: self._change_state( - "set_server_url", + "set_addressing", _StateAllocatedUninitialized, - _StateAllocatedUninitialized(actor_handle=self.actor_handle, server_url=server_url), + _StateAllocatedUninitialized(actor_handle=self.actor_handle, addr_info=addr_info), ) def mark_alive(self): self._change_state( "mark_alive", _StateAllocatedUninitialized, - _StateAllocatedAlive(actor_handle=self.actor_handle, server_url=self.server_url), + _StateAllocatedAlive(actor_handle=self.actor_handle, addr_info=self.addr_info), ) def mark_stopped(self): @@ -50,14 +50,14 @@ def actor_handle(self) -> ray.actor.ActorHandle: return self._state.actor_handle @property - def server_url(self) -> str: + def addr_info(self) -> AddrInfo: assert isinstance(self._state, _StateAllocatedBase) - assert self._state.server_url is not None, f"{self._state=}" - return self._state.server_url + assert self._state.addr_info is not None, f"{self._state=}" + return self._state.addr_info @property def api_client(self) -> SGLangApiClient: - return SGLangApiClient(server_url=self.server_url) + return SGLangApiClient(server_url=self.addr_info.server_url) @property def is_allocated(self) -> bool: @@ -80,6 +80,13 @@ def _change_state( logger.info(f"{debug_name} end new={self._state}") +class AddrInfo(BaseModel): + model_config = ConfigDict(extra="forbid", frozen=True) + + server_url: str + bootstrap_port: int | None = None + + # ------------------------- states ----------------------------- @@ -93,7 +100,7 @@ class _StateStopped(_StateBase): class _StateAllocatedBase(_StateBase): actor_handle: ray.actor.ActorHandle - server_url: str | None = None + addr_info: AddrInfo | None = None class _StateAllocatedUninitialized(_StateAllocatedBase): diff --git a/miles/ray/rollout/server_group.py b/miles/ray/rollout/server_group.py index b611edebec4..740ad1c7858 100644 --- a/miles/ray/rollout/server_group.py +++ b/miles/ray/rollout/server_group.py @@ -9,14 +9,15 @@ from sglang.srt.constants import GPU_MEMORY_TYPE_WEIGHTS from miles.backends.sglang_utils.sglang_engine import SGLangEngine, build_server_url +from miles.backends.sglang_utils.sglang_router_api_client import SGLangRouterApiClient, use_legacy_router_api from miles.ray.rollout.addr_allocator import ( PortCursors, allocate_rollout_engine_addr_and_ports_external, allocate_rollout_engine_addr_and_ports_normal, ) -from miles.ray.rollout.server_engine import ServerEngine +from miles.ray.rollout.server_engine import AddrInfo, ServerEngine from miles.ray.utils import NOSET_VISIBLE_DEVICES_ENV_VARS_LIST -from miles.utils import dumper_utils +from miles.utils import async_utils, dumper_utils logger = logging.getLogger(__name__) @@ -173,20 +174,57 @@ def start_engines( for index, _ in new_engines: engine_addr_and_ports = addr_and_ports[index] - self.all_engines[index - self.rank_offset].set_server_url( - build_server_url(host=engine_addr_and_ports["host"], port=engine_addr_and_ports["port"]) + self.all_engines[index - self.rank_offset].set_addressing( + AddrInfo( + server_url=build_server_url( + host=engine_addr_and_ports["host"], port=engine_addr_and_ports["port"] + ), + bootstrap_port=engine_addr_and_ports.get("disaggregation_bootstrap_port"), + ) ) - init_handles = [ - engine.init.remote( - **addr_and_ports[index], - router_ip=self.router_ip, - router_port=self.router_port, - ) - for index, engine in new_engines - ] + init_handles = [engine.init.remote(**addr_and_ports[index]) for index, engine in new_engines] return init_handles, new_engine_indices + async def register_workers(self, engine_indices: list[int]) -> None: + if self.args.rollout_external or not (self.router_ip and self.router_port): + return + await asyncio.gather( + *[ + self._router_api_client.add_worker( + worker_url=engine.addr_info.server_url, + worker_type=self.worker_type, + use_legacy_api=use_legacy_router_api(self.args), + bootstrap_port=engine.addr_info.bootstrap_port, + ) + for engine in self._primary_engines_of(engine_indices) + ] + ) + + async def unregister_workers(self, engine_indices: list[int]) -> None: + if self.args.rollout_external or not (self.router_ip and self.router_port): + return + await asyncio.gather( + *[ + self._router_api_client.remove_worker( + worker_url=engine.addr_info.server_url, + use_legacy_api=use_legacy_router_api(self.args), + ) + for engine in self._primary_engines_of(engine_indices) + ] + ) + + def _primary_engines_of(self, engine_indices: list[int]) -> list[ServerEngine]: + return [ + self.all_engines[index] + for index in engine_indices + if index % self.nodes_per_engine == 0 and self.all_engines[index].is_allocated + ] + + @property + def _router_api_client(self) -> SGLangRouterApiClient: + return SGLangRouterApiClient(router_url=f"http://{self.router_ip}:{self.router_port}") + # Called from InferenceController.stop_cell (main thread, async): deliberately non-async here # to avoid introducing two states like "stopping (but not stopped)" vs "stopped", since # single-thread async code will not yield without an await point @@ -194,6 +232,10 @@ def start_engines( # moving `shutdown` mainly to local code def stop_engines(self, engine_indices: list[int]): logger.info(f"Killing server {engine_indices=}...") + try: + async_utils.run(asyncio.wait_for(self.unregister_workers(engine_indices), timeout=_SHUTDOWN_TIMEOUT)) + except Exception as e: + logger.warning(f"Unregistering {engine_indices=} from the router failed, tearing down anyway (e: {e})") for i in engine_indices: engine = self.all_engines[i] if engine.is_allocated: @@ -242,6 +284,7 @@ async def recover(self, port_cursors: PortCursors, filter_indices: list[int] | N ) self.mark_alive(engine_indices=new_engine_indices) + await self.register_workers(new_engine_indices) def mark_alive(self, engine_indices: list[int]): for engine_index in engine_indices: diff --git a/tests/fast/backends/sglang_utils/test_sglang_engine_init_addressing.py b/tests/fast/backends/sglang_utils/test_sglang_engine_init_addressing.py new file mode 100644 index 00000000000..dff038a5830 --- /dev/null +++ b/tests/fast/backends/sglang_utils/test_sglang_engine_init_addressing.py @@ -0,0 +1,38 @@ +from argparse import Namespace +from unittest.mock import MagicMock, patch + +import pytest + +_MODULE = "miles.backends.sglang_utils.sglang_engine" + + +def test_init_brackets_bare_ipv6_inputs_and_keeps_the_allocated_port(): + """A bare v6 host makes host:port ambiguous, so init must bracket what it forwards and what + it serves on -- at the port the allocator handed it, not a fixed one.""" + pytest.importorskip("sglang") + from miles.backends.sglang_utils.sglang_engine import SGLangEngine + + engine = SGLangEngine.__new__(SGLangEngine) + engine.args = Namespace(env_report=None, rollout_external=False) + engine.rank = 0 + engine.worker_type = "regular" + engine.base_gpu_id = 0 + engine.sglang_overrides = {} + engine.num_gpus_per_engine = 1 + + forwarded: dict = {} + + def fake_compute_server_args(*args, **kwargs): + forwarded["dist_init_addr"] = args[2] + forwarded["host"] = args[4] + return {"node_rank": 0, "host": "fd00::2", "port": 31007, "disaggregation_bootstrap_port": None}, [] + + with ( + patch(f"{_MODULE}._compute_server_args", side_effect=fake_compute_server_args), + patch(f"{_MODULE}.ServerArgs"), + patch(f"{_MODULE}.launch_server_process", return_value=MagicMock(pid=4242)), + ): + engine.init(dist_init_addr="fd00::1:15003", port=31007, nccl_port=6000, host="fd00::2") + + assert forwarded == {"dist_init_addr": "[fd00::1]:15003", "host": "[fd00::2]"} + assert engine.server_url == "http://[fd00::2]:31007" diff --git a/tests/fast/backends/sglang_utils/test_sglang_engine_router_wiring.py b/tests/fast/backends/sglang_utils/test_sglang_engine_router_wiring.py deleted file mode 100644 index fb70f1bd25a..00000000000 --- a/tests/fast/backends/sglang_utils/test_sglang_engine_router_wiring.py +++ /dev/null @@ -1,272 +0,0 @@ -from argparse import Namespace -from unittest.mock import MagicMock, patch - -import pytest - -_MODULE = "miles.backends.sglang_utils.sglang_engine" - - -class _RecordingRouterApiClient: - def __init__(self, events: list[tuple[str, dict]] | None = None): - self.calls: list[tuple[str, dict]] = [] if events is None else events - - async def add_worker(self, **kwargs): - self.calls.append(("add_worker", kwargs)) - - async def remove_worker(self, **kwargs): - self.calls.append(("remove_worker", kwargs)) - - -@pytest.fixture(autouse=True) -def modern_router(monkeypatch): - pytest.importorskip("sglang_router") - import sglang_router - - monkeypatch.setattr(sglang_router, "__version__", "0.3.1") - - -def _make_engine( - *, - worker_type: str = "regular", - use_miles_router: bool = False, - rollout_external: bool = False, - events: list[tuple[str, dict]] | None = None, -): - pytest.importorskip("sglang") - from miles.backends.sglang_utils.sglang_engine import SGLangEngine - - engine = SGLangEngine.__new__(SGLangEngine) - engine.args = Namespace(use_miles_router=use_miles_router, rollout_external=rollout_external) - engine.rank = 0 - engine.worker_type = worker_type - engine.node_rank = 0 - engine.server_host = "10.0.0.1" - engine.server_port = 30000 - engine.server_url = "http://10.0.0.1:30000" - engine.router_ip = "10.0.0.9" - engine.router_port = 9000 - engine.router_api_client = _RecordingRouterApiClient(events) - engine.process = MagicMock(pid=4242) - return engine - - -def test_init_registers_the_engines_own_url_with_the_router(): - """The router must be told the url the engine actually serves on.""" - engine = _make_engine() - - with ( - patch(f"{_MODULE}.ServerArgs"), - patch(f"{_MODULE}.launch_server_process", return_value=MagicMock(pid=4242)), - ): - engine._init_normal({"disaggregation_bootstrap_port": None}) - - assert engine.router_api_client.calls == [ - ( - "add_worker", - { - "worker_url": "http://10.0.0.1:30000", - "worker_type": "regular", - "use_legacy_api": False, - "bootstrap_port": None, - }, - ) - ] - - -def test_init_passes_the_bootstrap_port_of_a_prefill_worker(): - """PD disaggregation needs the decode side to dial this port.""" - engine = _make_engine(worker_type="prefill") - - with ( - patch(f"{_MODULE}.ServerArgs"), - patch(f"{_MODULE}.launch_server_process", return_value=MagicMock(pid=4242)), - ): - engine._init_normal({"disaggregation_bootstrap_port": 8998}) - - assert len(engine.router_api_client.calls) == 1 - assert engine.router_api_client.calls[0][1]["bootstrap_port"] == 8998 - - -def test_init_skips_registration_on_non_node0_actors(): - """Only node 0 of a multi-node engine serves the router-visible endpoint.""" - engine = _make_engine() - engine.node_rank = 1 - - with ( - patch(f"{_MODULE}.ServerArgs"), - patch(f"{_MODULE}.launch_server_process", return_value=None), - ): - engine._init_normal({"disaggregation_bootstrap_port": None}) - - assert engine.router_api_client.calls == [] - - -def test_init_skips_registration_without_a_router(): - engine = _make_engine() - engine.router_ip = None - - with ( - patch(f"{_MODULE}.ServerArgs"), - patch(f"{_MODULE}.launch_server_process", return_value=MagicMock(pid=4242)), - ): - engine._init_normal({"disaggregation_bootstrap_port": None}) - - assert engine.router_api_client.calls == [] - - -def test_shutdown_unregisters_before_killing_the_server(): - """Killing first would leave the router routing to a dead worker.""" - events: list[tuple[str, dict]] = [] - engine = _make_engine(events=events) - - with patch(f"{_MODULE}.kill_process_tree", side_effect=lambda pid: events.append(("kill", {"pid": pid}))): - engine.shutdown() - - assert events == [ - ("remove_worker", {"worker_url": "http://10.0.0.1:30000", "use_legacy_api": False}), - ("kill", {"pid": 4242}), - ] - - -def test_shutdown_of_an_external_engine_touches_neither_router_nor_process(): - """An external engine is owned by someone else.""" - engine = _make_engine(rollout_external=True) - - with patch(f"{_MODULE}.kill_process_tree") as kill_mock: - engine.shutdown() - - assert engine.router_api_client.calls == [] - kill_mock.assert_not_called() - - -@pytest.mark.parametrize( - "server_host, expected_worker_url", - [("10.0.0.1", "http://10.0.0.1:30000"), ("[fd00::1]", "http://[fd00::1]:30000")], -) -def test_init_builds_the_router_client_and_worker_url_from_its_own_placement(server_host, expected_worker_url): - """init() derives the router url and the ipv6-safe worker url before any registration.""" - pytest.importorskip("sglang") - from miles.backends.sglang_utils.sglang_engine import SGLangEngine - - engine = SGLangEngine.__new__(SGLangEngine) - engine.args = Namespace(env_report=None, rollout_external=False, use_miles_router=False) - engine.rank = 0 - engine.worker_type = "regular" - engine.base_gpu_id = 0 - engine.sglang_overrides = {} - engine.num_gpus_per_engine = 1 - - recorder = _RecordingRouterApiClient() - constructor_kwargs: dict = {} - - def make_router_api_client(**kwargs): - constructor_kwargs.update(kwargs) - return recorder - - server_args_dict = {"node_rank": 0, "host": server_host, "port": 30000, "disaggregation_bootstrap_port": None} - with ( - patch(f"{_MODULE}._compute_server_args", return_value=(server_args_dict, [])), - patch(f"{_MODULE}.SGLangRouterApiClient", side_effect=make_router_api_client), - patch(f"{_MODULE}.ServerArgs"), - patch(f"{_MODULE}.launch_server_process", return_value=MagicMock(pid=4242)), - ): - engine.init( - dist_init_addr="10.0.0.1:5000", - port=30000, - nccl_port=6000, - host="10.0.0.1", - router_ip="10.0.0.9", - router_port=9000, - ) - - assert constructor_kwargs == {"router_url": "http://10.0.0.9:9000"} - assert recorder.calls == [ - ( - "add_worker", - { - "worker_url": expected_worker_url, - "worker_type": "regular", - "use_legacy_api": False, - "bootstrap_port": None, - }, - ) - ] - - -def test_init_brackets_bare_ipv6_inputs_and_keeps_the_allocated_port(): - """A bare v6 host makes host:port ambiguous, so init must bracket what it forwards and what - it serves on -- at the port the allocator handed it, not a fixed one.""" - pytest.importorskip("sglang") - from miles.backends.sglang_utils.sglang_engine import SGLangEngine - - engine = SGLangEngine.__new__(SGLangEngine) - engine.args = Namespace(env_report=None, rollout_external=False, use_miles_router=False) - engine.rank = 0 - engine.worker_type = "regular" - engine.base_gpu_id = 0 - engine.sglang_overrides = {} - engine.num_gpus_per_engine = 1 - - recorder = _RecordingRouterApiClient() - forwarded: dict = {} - - def fake_compute_server_args(*args, **kwargs): - forwarded["dist_init_addr"] = args[2] - forwarded["host"] = args[4] - return {"node_rank": 0, "host": "fd00::2", "port": 31007, "disaggregation_bootstrap_port": None}, [] - - with ( - patch(f"{_MODULE}._compute_server_args", side_effect=fake_compute_server_args), - patch(f"{_MODULE}.SGLangRouterApiClient", return_value=recorder), - patch(f"{_MODULE}.ServerArgs"), - patch(f"{_MODULE}.launch_server_process", return_value=MagicMock(pid=4242)), - ): - engine.init( - dist_init_addr="fd00::1:15003", - port=31007, - nccl_port=6000, - host="fd00::2", - router_ip="10.0.0.9", - router_port=9000, - ) - - assert forwarded == {"dist_init_addr": "[fd00::1]:15003", "host": "[fd00::2]"} - assert engine.server_url == "http://[fd00::2]:31007" - assert recorder.calls[0][1]["worker_url"] == "http://[fd00::2]:31007" - - -def test_init_and_shutdown_forward_the_legacy_api_decision(): - """--use-miles-router must reach both router calls, not just the helper.""" - events: list[tuple[str, dict]] = [] - engine = _make_engine(use_miles_router=True, events=events) - - with ( - patch(f"{_MODULE}.ServerArgs"), - patch(f"{_MODULE}.launch_server_process", return_value=MagicMock(pid=4242)), - patch(f"{_MODULE}.kill_process_tree"), - ): - engine._init_normal({"disaggregation_bootstrap_port": None}) - engine.shutdown() - - assert [kwargs["use_legacy_api"] for _name, kwargs in events] == [True, True] - - -@pytest.mark.parametrize( - "version, use_miles_router, expected", - [ - ("0.2.1", False, True), - ("0.2.2", False, False), - ("0.3.1", False, False), - ("0.3.1", True, True), - ], -) -def test_legacy_router_api_decision(version, use_miles_router, expected, monkeypatch): - """0.2.1 is the last version with the query-string API; --use-miles-router pins it too.""" - pytest.importorskip("sglang") - import sglang_router - - from miles.backends.sglang_utils.sglang_engine import _use_legacy_router_api - - monkeypatch.setattr(sglang_router, "__version__", version) - - assert _use_legacy_router_api(Namespace(use_miles_router=use_miles_router)) is expected diff --git a/tests/fast/backends/sglang_utils/test_sglang_router_api_client.py b/tests/fast/backends/sglang_utils/test_sglang_router_api_client.py index ec4b0a5ca94..ad37e1d7b6f 100644 --- a/tests/fast/backends/sglang_utils/test_sglang_router_api_client.py +++ b/tests/fast/backends/sglang_utils/test_sglang_router_api_client.py @@ -1,3 +1,5 @@ +from argparse import Namespace + import httpx import pytest @@ -247,3 +249,21 @@ async def get(self, url, **kwargs): await client.remove_worker(worker_url=WORKER_URL, use_legacy_api=False) assert "Failed to fetch workers list" in caplog.text + + +@pytest.mark.parametrize( + "version, use_miles_router, expected", + [ + ("0.2.1", False, True), + ("0.2.2", False, False), + ("0.3.1", False, False), + ("0.3.1", True, True), + ], +) +def test_legacy_router_api_decision(version, use_miles_router, expected, monkeypatch): + """0.2.1 is the last version with the query-string API; --use-miles-router pins it too.""" + from miles.backends.sglang_utils.sglang_router_api_client import use_legacy_router_api + + monkeypatch.setattr(sglang_router, "__version__", version) + + assert use_legacy_router_api(Namespace(use_miles_router=use_miles_router)) is expected diff --git a/tests/fast/ray/rollout/real_ray/test_fault_tolerance.py b/tests/fast/ray/rollout/real_ray/test_fault_tolerance.py index cc522c598ab..424a5d3f932 100644 --- a/tests/fast/ray/rollout/real_ray/test_fault_tolerance.py +++ b/tests/fast/ray/rollout/real_ray/test_fault_tolerance.py @@ -113,6 +113,42 @@ async def test_recover_default_filter_picks_all_dead_slots( finally: _kill_all(group) + async def test_recover_publishes_the_new_url_to_the_router( + self, + patched_sglang_engine, + placement_group_factory, + mock_engine_http_servers, + ): + """A recovered engine gets a fresh port, so the router must be told the new url.""" + from unittest.mock import patch + + from miles.ray.rollout.server_group import ServerGroup + + events: list[dict] = [] + + class _Recorder: + async def add_worker(self, **kwargs): + events.append(kwargs) + + async def remove_worker(self, **kwargs): + events.append(kwargs) + + pg = placement_group_factory(1) + group = _build_group(pg_tuple=pg, num_engines=1) + group.router_ip, group.router_port = "10.0.0.9", 9000 + _start(group) + ray.kill(group.all_engines[0].actor_handle) + group.all_engines[0].mark_stopped() + + try: + with patch.object(ServerGroup, "_router_api_client", property(lambda self: _Recorder())): + await group.recover(port_cursors=PortCursors.empty(), filter_indices=[0]) + + assert [event["worker_url"] for event in events] == [group.all_engines[0].addr_info.server_url] + assert group.all_engines[0].is_alive + finally: + _kill_all(group) + async def test_recover_with_offload_calls_release_then_resume( self, patched_sglang_engine, 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 0f09b1ebff9..9835e562fc6 100644 --- a/tests/fast/ray/rollout/real_ray/test_inference_controller.py +++ b/tests/fast/ray/rollout/real_ray/test_inference_controller.py @@ -12,11 +12,27 @@ from miles.ray.rollout.inference_controller import InferenceController +class _NoopRouterApiClient: + """The rollout process registers its engines for real; ``sglang_router_ip`` + here is a placeholder that keeps ``start_router`` short-circuited, and no + router listens on it.""" + + def __init__(self, router_url: str): + self.router_url = router_url + + async def add_worker(self, **kwargs): + return None + + async def remove_worker(self, **kwargs): + return None + + @pytest.fixture def patch_low_level(monkeypatch, mock_engine_http_servers): """Replace, in the test process: - ``SGLangEngine`` → ``MockSGLangEngine`` so created actors are mocks. - addr allocator → deterministic stub pointing at the mock http servers. + - ``SGLangRouterApiClient`` → no-op (no router runs at the placeholder address). - ``start_session_server`` → no-op (the production default touches network).""" import miles.ray.rollout.inference_controller as ictl import miles.ray.rollout.rollout_server as rsrv @@ -50,6 +66,7 @@ def _fake_alloc(*args, **kwargs): ) monkeypatch.setattr(sg, "allocate_rollout_engine_addr_and_ports_normal", _fake_alloc) + monkeypatch.setattr(sg, "SGLangRouterApiClient", _NoopRouterApiClient) monkeypatch.setattr(ictl, "start_session_server", lambda args: None) @@ -398,7 +415,7 @@ async def test_check_weights_targets_only_updatable_model( assert engine_result == {"mock": True} updatable_urls = { - engine.server_url + engine.addr_info.server_url for srv in controller.servers.values() if srv.update_weights for group in srv.server_groups @@ -406,7 +423,7 @@ async def test_check_weights_targets_only_updatable_model( if engine.is_allocated } frozen_urls = { - engine.server_url + engine.addr_info.server_url for srv in controller.servers.values() if not srv.update_weights for group in srv.server_groups diff --git a/tests/fast/ray/rollout/real_ray/test_rollout_server.py b/tests/fast/ray/rollout/real_ray/test_rollout_server.py index b7fdf5d4141..daa70894624 100644 --- a/tests/fast/ray/rollout/real_ray/test_rollout_server.py +++ b/tests/fast/ray/rollout/real_ray/test_rollout_server.py @@ -82,7 +82,7 @@ async def test_aggregates_across_groups_via_real_asyncio_gather( for rank in range(5) } for engine in all_engines: - server = url_to_server[engine.server_url] + server = url_to_server[engine.addr_info.server_url] payloads = server.payloads_of("/weights_checker") assert payloads == [{"action": "report", "allow_quant_error": False, "selector": "all"}] finally: diff --git a/tests/fast/ray/rollout/real_ray/test_server_group.py b/tests/fast/ray/rollout/real_ray/test_server_group.py index 811eb63a911..71b3c970802 100644 --- a/tests/fast/ray/rollout/real_ray/test_server_group.py +++ b/tests/fast/ray/rollout/real_ray/test_server_group.py @@ -81,7 +81,7 @@ def test_creates_real_actors_and_init_runs( init_kwargs = ray.get(e.actor_handle.get_init_kwargs.remote()) assert init_kwargs["host"] == "127.0.0.1" assert init_kwargs["port"] == mock_engine_http_servers.for_rank(i).port - assert e.server_url == mock_engine_http_servers.for_rank(i).url + assert e.addr_info.server_url == mock_engine_http_servers.for_rank(i).url # Cleanup: kill the actors we created. for e in group.all_engines: diff --git a/tests/fast/ray/rollout/test_server_cell.py b/tests/fast/ray/rollout/test_server_cell.py index b20b15aff73..51b939d612a 100644 --- a/tests/fast/ray/rollout/test_server_cell.py +++ b/tests/fast/ray/rollout/test_server_cell.py @@ -6,7 +6,7 @@ from miles.ray.rollout.rollout_server import RolloutServer from miles.ray.rollout.server_cell import get_cell_indexer_of_id_map -from miles.ray.rollout.server_engine import ServerEngine +from miles.ray.rollout.server_engine import AddrInfo, ServerEngine from miles.ray.rollout.server_group import ServerGroup @@ -21,7 +21,7 @@ def _build_servers( engines = [ServerEngine() for _ in range(engines_per_group)] for e in engines: e.mark_allocated_uninitialized(fake_actor_handle()) - e.set_server_url("http://127.0.0.1:30000") + e.set_addressing(AddrInfo(server_url="http://127.0.0.1:30000")) e.mark_alive() groups.append( ServerGroup( diff --git a/tests/fast/ray/rollout/test_server_engine.py b/tests/fast/ray/rollout/test_server_engine.py index 98c3f4518e9..140686fef11 100644 --- a/tests/fast/ray/rollout/test_server_engine.py +++ b/tests/fast/ray/rollout/test_server_engine.py @@ -3,7 +3,7 @@ import pytest import ray -from miles.ray.rollout.server_engine import ServerEngine +from miles.ray.rollout.server_engine import AddrInfo, ServerEngine def _fake_actor_handle() -> MagicMock: @@ -22,7 +22,7 @@ def test_api_client_is_unavailable_before_the_url_is_known(): def test_api_client_targets_the_assigned_url(): engine = ServerEngine() engine.mark_allocated_uninitialized(_fake_actor_handle()) - engine.set_server_url("http://10.0.0.1:30000") + engine.set_addressing(AddrInfo(server_url="http://10.0.0.1:30000")) assert engine.api_client.server_url == "http://10.0.0.1:30000" @@ -31,7 +31,7 @@ def test_mark_alive_keeps_the_url(): """Going alive keeps the assigned url.""" engine = ServerEngine() engine.mark_allocated_uninitialized(_fake_actor_handle()) - engine.set_server_url("http://10.0.0.1:30000") + engine.set_addressing(AddrInfo(server_url="http://10.0.0.1:30000")) engine.mark_alive() assert engine.is_alive @@ -51,7 +51,7 @@ def test_restart_replaces_the_url(): """A restarted engine takes the new url.""" engine = ServerEngine() engine.mark_allocated_uninitialized(_fake_actor_handle()) - engine.set_server_url("http://10.0.0.1:30000") + engine.set_addressing(AddrInfo(server_url="http://10.0.0.1:30000")) engine.mark_alive() engine.mark_stopped() @@ -60,6 +60,6 @@ def test_restart_replaces_the_url(): _ = engine.api_client engine.mark_allocated_uninitialized(_fake_actor_handle()) - engine.set_server_url("http://10.0.0.1:31000") + engine.set_addressing(AddrInfo(server_url="http://10.0.0.1:31000")) assert engine.api_client.server_url == "http://10.0.0.1:31000" diff --git a/tests/fast/ray/rollout/test_server_group_router_registration.py b/tests/fast/ray/rollout/test_server_group_router_registration.py new file mode 100644 index 00000000000..36f13aca337 --- /dev/null +++ b/tests/fast/ray/rollout/test_server_group_router_registration.py @@ -0,0 +1,211 @@ +from __future__ import annotations + +import asyncio +from unittest.mock import patch + +import pytest +from tests.fast.ray.rollout.conftest import fake_actor_handle, make_args + +from miles.ray.rollout.server_engine import AddrInfo, ServerEngine +from miles.ray.rollout.server_group import ServerGroup +from miles.utils import async_utils + +_MODULE = "miles.ray.rollout.server_group" + + +class _RecordingRouterApiClient: + def __init__(self, events: list[tuple[str, dict]], remove_worker_effect=None): + self._events = events + self._remove_worker_effect = remove_worker_effect + + async def add_worker(self, **kwargs): + self._events.append(("add_worker", kwargs)) + + async def remove_worker(self, **kwargs): + self._events.append(("remove_worker", kwargs)) + if self._remove_worker_effect is not None: + await self._remove_worker_effect() + + +def _build_group( + *, + events: list[tuple[str, dict]], + num_engines: int = 1, + num_gpus_per_engine: int = 1, + worker_type: str = "regular", + router_ip: str | None = "10.0.0.9", + router_port: int | None = 9000, + bootstrap_port: int | None = None, + use_miles_router: bool = False, + rollout_external: bool = False, + remove_worker_effect=None, +) -> ServerGroup: + args = make_args(num_gpus_per_node=8, use_miles_router=use_miles_router, rollout_external=rollout_external) + engines = [] + for index in range(num_engines): + engine = ServerEngine() + engine.mark_allocated_uninitialized(fake_actor_handle()) + engine.set_addressing( + AddrInfo(server_url=f"http://10.0.0.{index + 1}:3000{index}", bootstrap_port=bootstrap_port) + ) + engine.mark_alive() + engines.append(engine) + + group = ServerGroup( + args=args, + pg=(None, [], []), + all_engines=engines, + num_gpus_per_engine=num_gpus_per_engine, + has_new_engines=False, + worker_type=worker_type, + router_ip=router_ip, + router_port=router_port, + ) + group._recording_router_client = _RecordingRouterApiClient(events, remove_worker_effect=remove_worker_effect) + return group + + +def _with_recording_client(group: ServerGroup): + return patch.object(ServerGroup, "_router_api_client", property(lambda self: self._recording_router_client)) + + +async def test_registration_publishes_the_url_the_engine_actually_serves(): + """The router must be told the url the rollout process derived from the allocator.""" + events: list[tuple[str, dict]] = [] + group = _build_group(events=events) + + with _with_recording_client(group): + await group.register_workers([0]) + + assert events == [ + ( + "add_worker", + { + "worker_url": "http://10.0.0.1:30000", + "worker_type": "regular", + "use_legacy_api": False, + "bootstrap_port": None, + }, + ) + ] + + +async def test_registration_passes_the_bootstrap_port_of_a_prefill_worker(): + """PD disaggregation needs the decode side to dial this port.""" + events: list[tuple[str, dict]] = [] + group = _build_group(events=events, worker_type="prefill", bootstrap_port=8998) + + with _with_recording_client(group): + await group.register_workers([0]) + + assert events[0][1]["worker_type"] == "prefill" + assert events[0][1]["bootstrap_port"] == 8998 + + +async def test_registration_addresses_only_node0_of_a_multi_node_engine(): + """Only node 0 serves the router-visible endpoint.""" + events: list[tuple[str, dict]] = [] + group = _build_group(events=events, num_engines=2, num_gpus_per_engine=16) + + with _with_recording_client(group): + await group.register_workers([0, 1]) + + assert [kwargs["worker_url"] for _name, kwargs in events] == ["http://10.0.0.1:30000"] + + +@pytest.mark.parametrize("missing", [dict(router_ip=None), dict(router_port=None)]) +async def test_registration_is_skipped_without_a_router(missing): + events: list[tuple[str, dict]] = [] + group = _build_group(events=events, **missing) + + with _with_recording_client(group): + await group.register_workers([0]) + await group.unregister_workers([0]) + + assert events == [] + + +async def test_an_external_engine_is_never_registered_or_unregistered(): + """External engines are published by whoever runs them.""" + events: list[tuple[str, dict]] = [] + group = _build_group(events=events, rollout_external=True) + + with _with_recording_client(group): + await group.register_workers([0]) + await group.unregister_workers([0]) + + assert events == [] + + +def test_stop_engines_unregisters_before_killing_the_actor(): + """Killing first would leave the router routing to a dead worker.""" + events: list[tuple[str, dict]] = [] + group = _build_group(events=events) + + with ( + _with_recording_client(group), + patch(f"{_MODULE}.ray") as ray_mock, + ): + ray_mock.get.side_effect = lambda *args, **kwargs: events.append(("shutdown", {})) + ray_mock.kill.side_effect = lambda handle: events.append(("kill", {})) + group.stop_engines(engine_indices=[0]) + + assert [name for name, _kwargs in events] == ["remove_worker", "shutdown", "kill"] + assert events[0][1] == {"worker_url": "http://10.0.0.1:30000", "use_legacy_api": False} + + +def test_a_router_that_rejects_the_unregister_still_kills_the_actor(): + """Teardown is how a wedged engine is reclaimed, so a router error must not abort it.""" + + async def _reject(): + raise RuntimeError("router rejected the removal") + + events: list[tuple[str, dict]] = [] + group = _build_group(events=events, remove_worker_effect=_reject) + + with ( + _with_recording_client(group), + patch(f"{_MODULE}.ray") as ray_mock, + ): + ray_mock.get.side_effect = lambda *args, **kwargs: events.append(("shutdown", {})) + ray_mock.kill.side_effect = lambda handle: events.append(("kill", {})) + group.stop_engines(engine_indices=[0]) + + assert [name for name, _kwargs in events] == ["remove_worker", "shutdown", "kill"] + assert not group.all_engines[0].is_allocated + + +def test_a_router_that_never_answers_the_unregister_does_not_block_teardown(): + """The shared http client has no read timeout, so an unanswered removal would wedge teardown forever.""" + + async def _hang(): + await asyncio.sleep(3600) + + events: list[tuple[str, dict]] = [] + group = _build_group(events=events, remove_worker_effect=_hang) + + with ( + _with_recording_client(group), + patch(f"{_MODULE}._SHUTDOWN_TIMEOUT", 0.1), + patch(f"{_MODULE}.ray") as ray_mock, + ): + ray_mock.kill.side_effect = lambda handle: events.append(("kill", {})) + group.stop_engines(engine_indices=[0]) + + assert [name for name, _kwargs in events] == ["remove_worker", "kill"] + assert not group.all_engines[0].is_allocated + + +def test_use_miles_router_reaches_both_router_calls(): + """--use-miles-router pins the legacy query-string API on register and unregister alike.""" + events: list[tuple[str, dict]] = [] + group = _build_group(events=events, use_miles_router=True) + + with ( + _with_recording_client(group), + patch(f"{_MODULE}.ray"), + ): + async_utils.run(group.register_workers([0])) + group.stop_engines(engine_indices=[0]) + + assert [kwargs["use_legacy_api"] for _name, kwargs in events] == [True, True]