Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
64 changes: 64 additions & 0 deletions miles/utils/workers/worker_spec.py
Original file line number Diff line number Diff line change
@@ -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}
165 changes: 165 additions & 0 deletions tests/fast/utils/workers/test_worker_spec.py
Original file line number Diff line number Diff line change
@@ -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