diff --git a/miles/utils/workers/worker_spec.py b/miles/utils/workers/worker_spec.py new file mode 100644 index 00000000000..bd8be094770 --- /dev/null +++ b/miles/utils/workers/worker_spec.py @@ -0,0 +1,64 @@ +from collections.abc import Callable +from typing import Any, Literal + +from pydantic import model_validator + +from miles.utils.pydantic_utils import FrozenStrictBaseModel + +RPC_PORT_NAME = "rpc" +DEFAULT_RPC_PORT = 8000 + + +def _port_info_name(port_info: "PortInfo | dict") -> str: + return port_info["name"] if isinstance(port_info, dict) else port_info.name + + +class PortInfo(FrozenStrictBaseModel): + name: str + static_port: int + mode: Literal["per_worker", "master"] + allow_dynamic: bool + num_consecutive: int = 1 + offset_by_cell: bool = False + + @model_validator(mode="after") + def _reject_offsetting_a_dynamically_allocated_port(self) -> "PortInfo": + assert not ( + self.offset_by_cell and self.allow_dynamic + ), f"Port {self.name!r} cannot be offset by cell index: it is allocated dynamically" + return self + + +class SchedulingSpec(FrozenStrictBaseModel): + num_cells: int + num_workers_per_cell: int + num_gpus_per_worker: float + + +class BaseWorkerSpec(FrozenStrictBaseModel): + name: str + port_infos: list[PortInfo] + env_var: Callable[[], dict[str, str]] + scheduling: SchedulingSpec + + +class CommandWorkerSpec(BaseWorkerSpec): + launch_command: str + + +class ServeWorkerSpec(BaseWorkerSpec): + worker_class: str + ctor_kwargs: Callable[[], dict[str, Any]] + + @model_validator(mode="before") + @classmethod + def _inject_rpc_port(cls, values: dict) -> dict: + if "port_infos" not in values: + return values + + port_infos = list(values["port_infos"]) + if all(_port_info_name(port_info) != RPC_PORT_NAME for port_info in port_infos): + port_infos.append( + PortInfo(name=RPC_PORT_NAME, static_port=DEFAULT_RPC_PORT, mode="per_worker", allow_dynamic=True) + ) + return {**values, "port_infos": port_infos} diff --git a/tests/fast/utils/workers/test_worker_spec.py b/tests/fast/utils/workers/test_worker_spec.py new file mode 100644 index 00000000000..1164abadab5 --- /dev/null +++ b/tests/fast/utils/workers/test_worker_spec.py @@ -0,0 +1,165 @@ +import pytest +from pydantic import ValidationError + +from miles.utils.workers.serving import serve_inner +from miles.utils.workers.worker_spec import ( + DEFAULT_RPC_PORT, + RPC_PORT_NAME, + BaseWorkerSpec, + CommandWorkerSpec, + PortInfo, + SchedulingSpec, + ServeWorkerSpec, +) + + +def _make_port_info(**overrides) -> PortInfo: + kwargs = dict(name="http", static_port=8000, mode="per_worker", allow_dynamic=False) + kwargs.update(overrides) + return PortInfo(**kwargs) + + +def _make_base_kwargs(**overrides) -> dict: + kwargs = dict( + name="demo-worker", + port_infos=[_make_port_info()], + env_var=lambda: {"DEMO": "1"}, + scheduling=SchedulingSpec(num_cells=2, num_workers_per_cell=4, num_gpus_per_worker=0.4), + ) + kwargs.update(overrides) + return kwargs + + +class TestPortInfo: + def test_accepts_both_modes(self): + """Both per_worker and master are valid modes.""" + assert _make_port_info(mode="per_worker").mode == "per_worker" + assert _make_port_info(mode="master").mode == "master" + + def test_rejects_unknown_mode(self): + """An unknown mode literal is rejected.""" + with pytest.raises(ValidationError): + _make_port_info(mode="broadcast") + + def test_num_consecutive_defaults_to_one(self): + """A port reserves a single slot unless a block is requested.""" + assert _make_port_info().num_consecutive == 1 + assert _make_port_info(num_consecutive=32).num_consecutive == 32 + + def test_rejects_extra_field(self): + """Unknown fields are forbidden.""" + with pytest.raises(ValidationError): + _make_port_info(unknown_field=1) + + def test_is_frozen(self): + """Field assignment after construction is rejected.""" + port_info = _make_port_info() + with pytest.raises(ValidationError): + port_info.static_port = 9000 + + +class TestBaseWorkerSpec: + def test_constructs_and_exposes_fields(self): + """A spec keeps its name, ports, and scheduling as provided.""" + spec = BaseWorkerSpec(**_make_base_kwargs()) + assert spec.name == "demo-worker" + assert spec.port_infos[0].static_port == 8000 + assert spec.scheduling.num_cells == 2 + + def test_env_var_is_stored_uncalled(self): + """The env_var callable is stored as-is and only evaluated on demand.""" + calls = [] + + def env_var() -> dict[str, str]: + calls.append(1) + return {"A": "b"} + + spec = BaseWorkerSpec(**_make_base_kwargs(env_var=env_var)) + assert calls == [] + assert spec.env_var() == {"A": "b"} + + def test_rejects_extra_field(self): + """Unknown fields are forbidden.""" + with pytest.raises(ValidationError): + BaseWorkerSpec(**_make_base_kwargs(unknown_field=1)) + + def test_is_frozen(self): + """Field assignment after construction is rejected.""" + spec = BaseWorkerSpec(**_make_base_kwargs()) + with pytest.raises(ValidationError): + spec.name = "other" + + +class TestCommandWorkerSpec: + def test_constructs_with_launch_command(self): + """A command spec carries the launch command besides base fields.""" + spec = CommandWorkerSpec(**_make_base_kwargs(), launch_command="python -m sglang.launch_server") + assert spec.launch_command == "python -m sglang.launch_server" + assert isinstance(spec, BaseWorkerSpec) + + +class TestServeWorkerSpec: + def test_constructs_with_worker_class(self): + """A serve spec carries the worker class path besides base fields.""" + spec = ServeWorkerSpec( + **_make_base_kwargs(), + worker_class="miles.ray.rollout.inference_controller.InferenceController", + ctor_kwargs=lambda: {}, + ) + assert spec.worker_class == "miles.ray.rollout.inference_controller.InferenceController" + assert isinstance(spec, BaseWorkerSpec) + + def test_ctor_kwargs_is_stored_uncalled(self): + """The ctor_kwargs callable is stored as-is and only evaluated on demand.""" + calls = [] + + def ctor_kwargs() -> dict: + calls.append(1) + return {"x": 1} + + spec = ServeWorkerSpec( + **_make_base_kwargs(), + worker_class="miles.demo.Worker", + ctor_kwargs=ctor_kwargs, + ) + assert calls == [] + assert spec.ctor_kwargs() == {"x": 1} + + +class TestServeWorkerSpecRpcPortInjection: + def _make_spec(self, **overrides) -> ServeWorkerSpec: + return ServeWorkerSpec( + **_make_base_kwargs(**overrides), + worker_class="miles.demo.Worker", + ctor_kwargs=lambda: {}, + ) + + def test_rpc_port_is_injected_by_default(self): + """Every serve worker automatically exposes an rpc port.""" + spec = self._make_spec() + (rpc,) = [port_info for port_info in spec.port_infos if port_info.name == RPC_PORT_NAME] + assert rpc.static_port == DEFAULT_RPC_PORT + assert rpc.mode == "per_worker" + assert rpc.allow_dynamic is True + + def test_injection_keeps_declared_ports(self): + """The injected rpc port is appended after the declared ports.""" + spec = self._make_spec() + assert [port_info.name for port_info in spec.port_infos] == ["http", RPC_PORT_NAME] + + def test_explicit_rpc_port_is_not_duplicated(self): + """An explicitly declared rpc port wins over the injected default.""" + explicit = PortInfo(name=RPC_PORT_NAME, static_port=9999, mode="per_worker", allow_dynamic=False) + spec = self._make_spec(port_infos=[explicit]) + assert spec.port_infos == [explicit] + + def test_base_and_command_specs_get_no_rpc_port(self): + """Only serve workers run the rpc server, so only they get the port.""" + base = BaseWorkerSpec(**_make_base_kwargs()) + command = CommandWorkerSpec(**_make_base_kwargs(), launch_command="sleep 1") + assert RPC_PORT_NAME not in [port_info.name for port_info in base.port_infos] + assert RPC_PORT_NAME not in [port_info.name for port_info in command.port_infos] + + def test_the_injected_port_is_the_one_the_serve_entrypoint_binds_by_default(self): + """A spec advertising a port its own process does not bind leaves every caller talking to nothing.""" + assert DEFAULT_RPC_PORT == serve_inner.DEFAULT_PORT