Skip to content
Merged
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
17 changes: 7 additions & 10 deletions miles/backends/sglang_utils/sglang_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,13 +48,7 @@ def build_server_url(host: str, port: int) -> str:
return f"http://{format_v6_uri(host)}:{port}"


@dataclasses.dataclass(frozen=True)
class EngineLaunchPlan:
cmd: str
api_key: str | None


def compute_engine_launch_plan(
def compute_engine_launch_cmd(
args,
*,
node_rank: int,
Expand All @@ -63,7 +57,7 @@ def compute_engine_launch_plan(
sglang_overrides: dict,
num_gpus_per_engine: int,
addr_and_ports: dict,
) -> EngineLaunchPlan:
) -> str:
server_args_dict = _compute_server_args(
args,
node_rank=node_rank,
Expand All @@ -80,8 +74,11 @@ def compute_engine_launch_plan(
)

launch_args = {**server_args_dict, "host": server_args_dict["host"].strip("[]")}
cmd = shlex.join([sys.executable, "-m", "sglang.launch_server", *server_args_to_argv(launch_args)])
return EngineLaunchPlan(cmd=cmd, api_key=server_args_dict.get("api_key"))
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(
Expand Down
15 changes: 10 additions & 5 deletions miles/ray/rollout/server_cell.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,12 @@
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_engine_launch_plan, format_v6_uri
from miles.backends.sglang_utils.sglang_engine import (
build_server_url,
compute_api_key,
compute_engine_launch_cmd,
format_v6_uri,
)
from miles.backends.sglang_utils.sglang_router_api_client import SGLangRouterApiClient, use_legacy_router_api
from miles.ray.rollout.cell_state import (
AddrInfo,
Expand Down Expand Up @@ -154,8 +159,8 @@ async def start_engines(self, port_allocator: PortAllocator) -> None:
]
)

plans = {
rank: compute_engine_launch_plan(
launch_cmds = {
rank: compute_engine_launch_cmd(
self.args,
node_rank=local_index,
worker_type=self.worker_type,
Expand All @@ -169,14 +174,14 @@ async def start_engines(self, port_allocator: PortAllocator) -> None:

await asyncio.gather(
*[
actor.run.remote(cmd=plans[global_rank].cmd, envs={})
actor.run.remote(cmd=launch_cmds[global_rank], envs={})
for global_rank, actor in zip(global_ranks, actor_handles, strict=True)
]
)

await wait_server_healthy(
server_url=self.addr_info.server_url,
api_key=plans[global_ranks[0]].api_key,
api_key=compute_api_key(self.args, sglang_overrides=self.sglang_overrides),
is_process_alive=functools.partial(_engine_actor_is_alive, self.primary_actor_handle),
)

Expand Down
1 change: 1 addition & 0 deletions tests/fast/backends/sglang_utils/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -49,6 +49,7 @@ def make_engine_args(**overrides: Any) -> Namespace:
use_rollout_indexer_replay=False,
fp16=False,
lora_rank=0,
sglang_api_key=None,
lora_adapter_path=None,
multi_lora=False,
colocate=False,
Expand Down
49 changes: 30 additions & 19 deletions tests/fast/backends/sglang_utils/test_sglang_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,10 +9,10 @@
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_engine_launch_plan
from miles.backends.sglang_utils.sglang_engine import compute_api_key, compute_engine_launch_cmd


def _plan(*, worker_type: str = "regular", args=None, addr_overrides: dict | None = None, **kwargs):
def _cmd(*, worker_type: str = "regular", args=None, addr_overrides: dict | None = None, **kwargs) -> str:
addr_and_ports = dict(
host="10.0.0.1",
port=30000,
Expand All @@ -21,7 +21,7 @@ def _plan(*, worker_type: str = "regular", args=None, addr_overrides: dict | Non
dist_init_addr="10.0.0.1:20000",
)
addr_and_ports.update(addr_overrides or {})
return compute_engine_launch_plan(
return compute_engine_launch_cmd(
args or make_engine_args(),
node_rank=0,
worker_type=worker_type,
Expand All @@ -33,11 +33,10 @@ def _plan(*, worker_type: str = "regular", args=None, addr_overrides: dict | Non
)


class TestComputeEngineLaunchPlan:
class TestComputeEngineLaunchCmd:
def test_the_command_launches_sglang_with_the_allocated_addressing(self):
"""The plan renders one launch_server command carrying the addr map."""
plan = _plan()
tokens = shlex.split(plan.cmd)
"""The rendered launch_server command carries the addr map."""
tokens = shlex.split(_cmd())
assert tokens[:3] == [sys.executable, "-m", "sglang.launch_server"]
parsed = parse_server_args_argv(tokens[3:])
assert parsed.host == "10.0.0.1" and parsed.port == 30000
Expand All @@ -53,24 +52,36 @@ def test_every_plan_picks_a_fresh_random_seed(self):

def test_a_bracketed_v6_host_is_stripped_for_the_server_but_kept_in_dist_addr(self):
"""sglang binds a bare v6 host while the rendezvous addr stays bracketed."""
plan = _plan(addr_overrides=dict(host="[fd00::2]", port=31007, dist_init_addr="[fd00::1]:15003"))
parsed = parse_server_args_argv(shlex.split(plan.cmd)[3:])
cmd = _cmd(addr_overrides=dict(host="[fd00::2]", port=31007, dist_init_addr="[fd00::1]:15003"))
parsed = parse_server_args_argv(shlex.split(cmd)[3:])
assert parsed.host == "fd00::2"
assert parsed.dist_init_addr == "[fd00::1]:15003"

def test_a_prefill_plan_carries_the_bootstrap_port(self):
def test_a_prefill_command_carries_the_bootstrap_port(self):
"""PD-disaggregation prefill flags survive into the command."""
plan = _plan(worker_type="prefill", addr_overrides=dict(disaggregation_bootstrap_port=20090))
parsed = parse_server_args_argv(shlex.split(plan.cmd)[3:])
cmd = _cmd(worker_type="prefill", addr_overrides=dict(disaggregation_bootstrap_port=20090))
parsed = parse_server_args_argv(shlex.split(cmd)[3:])
assert parsed.disaggregation_mode == "prefill"
assert parsed.disaggregation_bootstrap_port == 20090

def test_the_plan_exposes_the_api_key_for_the_health_wait(self):
"""The driver-side health wait needs the same api key the server got."""
args = make_engine_args()
args.sglang_api_key = "secret"
assert _plan(args=args).api_key == "secret"
def test_the_command_carries_the_api_key_from_args(self):
"""--sglang-api-key reaches the server through the generic passthrough."""
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_the_plan_has_no_api_key_when_the_server_has_none(self):
def test_no_key_anywhere_means_the_health_wait_sends_none(self):
"""No key configured means the health wait sends none."""
assert _plan().api_key is None
assert compute_api_key(make_engine_args(), sglang_overrides={}) is None
Loading