From 5bc48512ea6532d3d51bfb6e674517db4c97f98e Mon Sep 17 00:00:00 2001 From: fzyzcjy Date: Sun, 2 Aug 2026 20:29:51 +0800 Subject: [PATCH] Trim ServerCellMetadata down to the fields a cell actually needs --- miles/backends/sglang_utils/sglang_config.py | 5 +++ miles/backends/sglang_utils/sglang_engine.py | 4 --- miles/dashboard/hooks.py | 16 ++++------ miles/ray/rollout/rollout_server.py | 13 +++++--- miles/ray/rollout/server_cell.py | 25 ++++----------- miles/utils/workers/naming.py | 7 ++++ .../sglang_utils/test_sglang_config.py | 32 +++++++++++++++++++ .../sglang_utils/test_sglang_engine.py | 18 +---------- 8 files changed, 65 insertions(+), 55 deletions(-) diff --git a/miles/backends/sglang_utils/sglang_config.py b/miles/backends/sglang_utils/sglang_config.py index fb41b1c0a74..4eaa46b9c3e 100644 --- a/miles/backends/sglang_utils/sglang_config.py +++ b/miles/backends/sglang_utils/sglang_config.py @@ -161,6 +161,11 @@ def resolve( default_model_path: str, gpu_offset_cursor: "_MutableBox", ) -> "ServerGroupConfig": + assert not ({"host", "port"} & set(raw.overrides)), ( + f"sglang_overrides must not override host/port ({raw.overrides=}): the rollout process derives " + f"each engine's url from the addr allocator, so an override would make it talk to the wrong endpoint" + ) + rollout_pg_offset = _compute_rollout_offset(args) megatron_num_gpus = _compute_megatron_num_gpus(args) diff --git a/miles/backends/sglang_utils/sglang_engine.py b/miles/backends/sglang_utils/sglang_engine.py index 910d3d2160f..c5c718ed32a 100644 --- a/miles/backends/sglang_utils/sglang_engine.py +++ b/miles/backends/sglang_utils/sglang_engine.py @@ -82,10 +82,6 @@ def compute_engine_launch_cmd( return shlex.join([sys.executable, "-m", "sglang.launch_server", *server_args_to_argv(launch_args)]) -def compute_api_key(args, *, sglang_overrides: dict) -> str | None: - return sglang_overrides.get("api_key", args.sglang_api_key) - - def _compute_server_args( args, *, diff --git a/miles/dashboard/hooks.py b/miles/dashboard/hooks.py index 424768a083b..cdf48456c0f 100644 --- a/miles/dashboard/hooks.py +++ b/miles/dashboard/hooks.py @@ -27,6 +27,7 @@ ) from miles.utils.lifecycle import TrajectoryLifecycle from miles.utils.timer import Timer +from miles.utils.workers.naming import parse_worker_name logger = logging.getLogger(__name__) @@ -380,19 +381,14 @@ def _alive_engine_cells(servers) -> list: def _collect_worker_infos(cells) -> list[list]: - from miles.ray.specs.inference import compute_engine_pool from miles.utils.workers.ray_worker_manager import RayWorkerManager manager_handle = RayWorkerManager.get_handle() - return _ray_get( - [ - manager_handle.get_worker_infos.remote( - pool=compute_engine_pool(model_idx=cell.meta.model_idx, group_index=cell.meta.group_index), - cell_index=cell.meta.cell_index, - ) - for cell in cells - ] - ) + futures = [] + for cell in cells: + pool, cell_index, _ = parse_worker_name(cell.meta.worker_name) + futures.append(manager_handle.get_worker_infos.remote(pool=pool, cell_index=cell_index)) + return _ray_get(futures) def _compute_engine_infos(cells, worker_infos_per_cell) -> list[EngineInfo]: diff --git a/miles/ray/rollout/rollout_server.py b/miles/ray/rollout/rollout_server.py index ce7b7b23c22..f788a80aefc 100644 --- a/miles/ray/rollout/rollout_server.py +++ b/miles/ray/rollout/rollout_server.py @@ -8,6 +8,8 @@ from miles.backends.sglang_utils.sglang_router_api_client import SGLangRouterApiClient from miles.ray.rollout.router_manager import wait_router_ready from miles.ray.rollout.server_cell import ServerCell, ServerCellMetadata, compute_nodes_per_engine +from miles.ray.specs.inference import compute_engine_pool +from miles.utils.workers.naming import compute_worker_name logger = logging.getLogger(__name__) @@ -63,17 +65,18 @@ async def start_rollout_servers(args) -> dict[str, "RolloutServer"]: ) for cell_start in range(0, num_engines, nodes_per_engine): cell_id = format_cell_id(server_id=model_cfg.name, index=cell_count) + cell_index = cell_start // nodes_per_engine + pool = compute_engine_pool(model_idx=model_idx, group_index=group_index) + worker_name = compute_worker_name(pool=pool, cell_index=cell_index) cell_meta = ServerCellMetadata( + model_id=model_cfg.name, worker_type=group_cfg.worker_type, cell_id=cell_id, num_gpus_per_engine=gpus_per_engine, gpu_offset=group_cfg.gpu_offset + cell_start * num_gpu_per_engine_local, - sglang_overrides=group_cfg.overrides, - model_idx=model_idx, - group_index=group_index, - cell_index=cell_start // nodes_per_engine, + sglang_api_key=group_cfg.overrides.get("api_key", args.sglang_api_key), + worker_name=worker_name, needs_offload=group_cfg.needs_offload, - model_path=group_cfg.model_path, update_weights=model_cfg.update_weights, ) cell_count += 1 diff --git a/miles/ray/rollout/server_cell.py b/miles/ray/rollout/server_cell.py index dcb0d1ac752..a5f8a7c6ad2 100644 --- a/miles/ray/rollout/server_cell.py +++ b/miles/ray/rollout/server_cell.py @@ -7,7 +7,7 @@ from sglang.srt.constants import GPU_MEMORY_TYPE_WEIGHTS from miles.backends.sglang_utils.sglang_api_client import SGLangApiClient, wait_server_healthy -from miles.backends.sglang_utils.sglang_engine import build_server_url, compute_api_key +from miles.backends.sglang_utils.sglang_engine import build_server_url from miles.backends.sglang_utils.sglang_router_api_client import SGLangRouterApiClient, use_legacy_router_api from miles.ray.rollout.cell_state import ( AddrInfo, @@ -16,9 +16,7 @@ StateAllocatedBase, StateAllocatedUninitialized, ) -from miles.ray.specs.inference import compute_engine_pool from miles.utils.pydantic_utils import FrozenStrictBaseModel -from miles.utils.workers.naming import compute_worker_name from miles.utils.workers.worker_provider.base import BaseWorkerProvider from miles.utils.workers.worker_provider.ray import RayWorkerProvider @@ -28,16 +26,14 @@ class ServerCellMetadata(FrozenStrictBaseModel): + model_id: str worker_type: Literal["regular", "prefill", "decode"] cell_id: str num_gpus_per_engine: int gpu_offset: int - sglang_overrides: dict - model_idx: int - group_index: int - cell_index: int + sglang_api_key: str | None + worker_name: str needs_offload: bool - model_path: str | None update_weights: bool @@ -65,23 +61,14 @@ def addr_info(self) -> AddrInfo: def api_client(self) -> SGLangApiClient: return SGLangApiClient(server_url=self.addr_info.server_url) - @property - def _pool_id(self) -> str: - return compute_engine_pool(model_idx=self.meta.model_idx, group_index=self.meta.group_index) - async def _add_raw(self) -> None: - assert not ({"host", "port"} & set(self.meta.sglang_overrides)), ( - f"sglang_overrides must not override host/port ({self.meta.sglang_overrides=}): the rollout process derives " - f"each engine's url from the addr allocator, so an override would make it talk to the wrong endpoint" - ) if self.args.rollout_external: raise NotImplementedError( "external rollout address allocation was removed and a new implementation is coming" ) provider: BaseWorkerProvider = RayWorkerProvider.create() # TODO inject instance - worker_name = compute_worker_name(pool=self._pool_id, cell_index=self.meta.cell_index) - master_addrs = await provider.get_addrs(worker_name=worker_name) + master_addrs = await provider.get_addrs(worker_name=self.meta.worker_name) primary = master_addrs["primary"] disaggregation_bootstrap = master_addrs.get("disaggregation_bootstrap") # TODO simplify (remove) later @@ -94,7 +81,7 @@ async def _add_raw(self) -> None: await wait_server_healthy( server_url=self.addr_info.server_url, - api_key=compute_api_key(self.args, sglang_overrides=self.meta.sglang_overrides), + api_key=self.meta.sglang_api_key, ) async def add(self, router_api_client: SGLangRouterApiClient, recover: bool = False) -> None: diff --git a/miles/utils/workers/naming.py b/miles/utils/workers/naming.py index dcf83341eb7..6ab0936a4db 100644 --- a/miles/utils/workers/naming.py +++ b/miles/utils/workers/naming.py @@ -1,2 +1,9 @@ +# TODO refactor & move later def compute_worker_name(*, pool_id: str, cell_index: int = 0, worker_in_cell_index: int = 0) -> str: return f"{pool_id}-{cell_index}-{worker_in_cell_index}" + + +# TODO refactor & move later +def parse_worker_name(worker_name: str) -> tuple[str, int, int]: + pool_id, cell_index, worker_in_cell_index = worker_name.rsplit("-", maxsplit=2) + return pool_id, int(cell_index), int(worker_in_cell_index) diff --git a/tests/fast/backends/sglang_utils/test_sglang_config.py b/tests/fast/backends/sglang_utils/test_sglang_config.py index d36578adac1..ebcb4cb8892 100644 --- a/tests/fast/backends/sglang_utils/test_sglang_config.py +++ b/tests/fast/backends/sglang_utils/test_sglang_config.py @@ -288,3 +288,35 @@ def test_an_explicit_memory_saver_override_wins(self, tmp_path): group = cfg.models[0].server_groups[0] assert group.needs_offload is False assert group.overrides["enable_memory_saver"] is True + + +class TestHostPortOverrideRejection: + def test_a_port_override_is_rejected_at_resolve_time(self, tmp_path): + """Overriding the allocator-owned port must fail fast instead of desyncing engine and controller.""" + with pytest.raises(AssertionError, match="must not override host/port"): + _resolve_yaml( + tmp_path, + "sglang:\n" + " - name: actor\n" + " server_groups:\n" + " - worker_type: regular\n" + " num_gpus: 8\n" + " overrides:\n" + " port: 12345\n", + rollout_num_gpus=8, + ) + + def test_a_host_override_is_rejected_at_resolve_time(self, tmp_path): + """Overriding the allocator-owned host must fail fast as well.""" + with pytest.raises(AssertionError, match="must not override host/port"): + _resolve_yaml( + tmp_path, + "sglang:\n" + " - name: actor\n" + " server_groups:\n" + " - worker_type: regular\n" + " num_gpus: 8\n" + " overrides:\n" + " host: 10.0.0.1\n", + rollout_num_gpus=8, + ) diff --git a/tests/fast/backends/sglang_utils/test_sglang_engine.py b/tests/fast/backends/sglang_utils/test_sglang_engine.py index 8805b883b0f..0c8b81a7c60 100644 --- a/tests/fast/backends/sglang_utils/test_sglang_engine.py +++ b/tests/fast/backends/sglang_utils/test_sglang_engine.py @@ -9,7 +9,7 @@ pytest.importorskip("sglang") from miles.backends.sglang_utils.server_args_utils import parse_server_args_argv -from miles.backends.sglang_utils.sglang_engine import compute_api_key, compute_engine_launch_cmd +from miles.backends.sglang_utils.sglang_engine import compute_engine_launch_cmd def _cmd(*, worker_type: str = "regular", args=None, addr_overrides: dict | None = None, **kwargs) -> str: @@ -75,19 +75,3 @@ def test_the_command_carries_the_api_key_from_args(self): cmd = _cmd(args=make_engine_args(sglang_api_key="secret")) parsed = parse_server_args_argv(shlex.split(cmd)[3:]) assert parsed.api_key == "secret" - - -class TestComputeApiKey: - def test_the_args_key_is_used_when_no_override_exists(self): - """The health wait needs the same key the generic passthrough gave the server.""" - args = make_engine_args(sglang_api_key="secret") - assert compute_api_key(args, sglang_overrides={}) == "secret" - - def test_an_override_key_wins_over_the_args_key(self): - """Overrides beat args exactly like they do in the rendered command.""" - args = make_engine_args(sglang_api_key="from-args") - assert compute_api_key(args, sglang_overrides={"api_key": "from-override"}) == "from-override" - - def test_no_key_anywhere_means_the_health_wait_sends_none(self): - """No key configured means the health wait sends none.""" - assert compute_api_key(make_engine_args(), sglang_overrides={}) is None