sgl-router: experimental Rust HTTP router for SGLang worker pools (#25851)

Co-authored-by: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
Kangyan-Zhou
2026-05-25 15:34:05 +08:00
committed by GitHub
co-authored by Claude Opus 4.7
parent aae04b1241
commit 6e8fe176be
131 changed files with 28623 additions and 55 deletions
@@ -0,0 +1,404 @@
"""Minimal sgl-router Gateway class — adapted from SMG's e2e_test/infra/gateway.py.
Differences from SMG:
- SMG drives a Python launcher (`python3 -m sglang_router.launch_router`)
with worker URLs on the CLI.
- sgl-router uses a Rust binary (`experimental/sgl-router/target/release/sgl-router`)
with a TOML config file. Worker discovery is config-file-based; this
Gateway writes a TOML to a tempfile and execs the binary with
`--config <tempfile>`.
Supported lifecycles:
- Regular mode: one model, N worker URLs, single policy.
- PD mode: one model, prefill_workers + decode_workers (lists of URLs),
discovery emits separate `WorkerMode::Prefill` / `WorkerMode::Decode`
entries. The router resolves PD pool isolation at request time.
Use as a context manager:
with Gateway() as gw:
gw.start_regular(model_path="...", worker_urls=[...])
resp = httpx.post(f"{gw.base_url}/v1/chat/completions", json=...)
or pytest fixture style (see e2e_test/conftest.py).
"""
from __future__ import annotations
import logging
import os
import signal
import socket
import subprocess
import tempfile
import time
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any
import httpx
logger = logging.getLogger(__name__)
# Repo-relative path to the release binary. Set ``SGL_ROUTER_BINARY`` to
# override (e.g. a debug build, or a non-default ``CARGO_TARGET_DIR``).
# This file is at `experimental/sgl-router/tests/e2e/infra/gateway.py`,
# so four `.parent` hops to reach the sgl-router workspace root
# (infra → e2e → tests → sgl-router). Cargo lands the binary at
# `experimental/sgl-router/target/release/sgl-router`. A previous
# version used three hops and pointed at `tests/target/`, which
# would have broken any test that actually launches the router via
# this helper.
DEFAULT_BINARY = (
Path(__file__).resolve().parent.parent.parent.parent
/ "target"
/ "release"
/ "sgl-router"
)
def _get_open_port() -> int:
"""Reserve an ephemeral TCP port in [20000, 55535].
The router itself doesn't have the ``port + 10000`` gRPC-derivation
constraint that SGLang's launch_server does, but we cap the range
anyway so the e2e helpers behave consistently across components.
"""
for _ in range(50):
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
s.bind(("127.0.0.1", 0))
port = s.getsockname()[1]
if 20000 <= port <= 55535:
return port
raise RuntimeError(
"could not allocate an ephemeral port in [20000, 55535] after 50 tries"
)
def _resolve_tokenizer_path(tokenizer_path: str) -> str:
"""Resolve a HuggingFace repo ID to a local ``tokenizer.json`` path.
sgl-router's tokenizer loader treats the input as a filesystem path and
inspects its extension; a bare HF id like ``Qwen/Qwen3-0.6B`` looks
like a file with extension ``.6B`` and is rejected. When the HF Hub
cache already has the tokenizer, point the loader at the on-disk
``tokenizer.json`` directly. Pass paths/URLs through unchanged.
"""
p = Path(tokenizer_path)
if p.exists():
return str(p)
try:
from huggingface_hub import try_to_load_from_cache # type: ignore[import]
cached = try_to_load_from_cache(tokenizer_path, "tokenizer.json")
if cached and Path(cached).is_file():
return str(cached)
except Exception: # noqa: BLE001
pass
return tokenizer_path
@dataclass
class WorkerInfo:
"""Worker visible to the gateway via ``/v1/models``-style introspection.
Mirrors SMG's WorkerInfo shape so test code reads the same. sgl-router
does not currently surface a `/v1/workers` admin API — this is a
placeholder for a future admin surface; current tests scrape
`/metrics` for per-worker observability instead.
"""
id: str
url: str
model: str | None = None
status: str = "unknown"
metadata: dict[str, Any] = field(default_factory=dict)
class Gateway:
"""Lifecycle-managed sgl-router instance for e2e tests.
Not thread-safe; assume one Gateway per test (or per fixture scope).
"""
def __init__(
self,
host: str = "127.0.0.1",
port: int | None = None,
binary: Path | None = None,
proxy_request_timeout_secs: int | None = None,
stale_request_timeout_secs: int | None = None,
):
self.host = host
self.port = port or _get_open_port()
self.base_url = f"http://{self.host}:{self.port}"
# Resolve binary from env override, explicit arg, or repo default.
env_binary = os.environ.get("SGL_ROUTER_BINARY")
if binary is not None:
self.binary = Path(binary)
elif env_binary:
self.binary = Path(env_binary)
else:
self.binary = DEFAULT_BINARY
# Test-side overrides for the router's tunables. Both default to
# `None`, in which case the router uses its production defaults
# (60 s proxy timeout, 300 s stale-request timeout). Tests set
# these short so per-request failures and stale-request expiry
# surface within the test's wall-time budget.
self.proxy_request_timeout_secs = proxy_request_timeout_secs
self.stale_request_timeout_secs = stale_request_timeout_secs
self.process: subprocess.Popen | None = None
self._config_path: Path | None = None
self._started: bool = False
# Track child workers we spawned so __exit__ can tear them down.
self._owned_workers: list[subprocess.Popen] = []
# ----- context manager -------------------------------------------------
def __enter__(self) -> "Gateway":
return self
def __exit__(self, *exc) -> None:
self.shutdown()
# ----- start ----------------------------------------------------------
def start_regular(
self,
*,
model_id: str,
tokenizer_path: str,
worker_urls: list[str],
policy: str = "round_robin",
extra_models: list[dict] | None = None,
timeout: float = 60.0,
) -> None:
"""Start the router in regular (non-PD) mode.
Args:
model_id: The model identifier the router will dispatch under.
tokenizer_path: Path or HF ID for the tokenizer the router uses
for cache-aware tokenization.
worker_urls: URLs of already-running ``sglang.launch_server``
instances. The router uses ``static_urls`` discovery;
each worker's mode (plain) and any disaggregation
metadata are learned from ``/server_info``.
policy: Policy kind — ``round_robin``, ``random``, ``power_of_two``,
or ``cache_aware_zmq``.
timeout: How long to wait for ``/readyz`` before giving up.
"""
self._launch(
self._build_config(
model_id=model_id,
tokenizer_path=tokenizer_path,
urls=list(worker_urls),
policy=policy,
extra_models=extra_models or [],
),
timeout=timeout,
)
def start_pd(
self,
*,
model_id: str,
tokenizer_path: str,
prefill_urls: list[str],
decode_urls: list[str],
policy: str = "round_robin",
timeout: float = 60.0,
) -> None:
"""Start the router in PD-disaggregated mode.
All prefill + decode URLs go into one ``static_urls`` list. The
router seeds each worker as ``WorkerMode::Plain`` and the
manager's ``/server_info`` introspect step overrides mode +
``bootstrap_port`` from the worker's self-disclosure. Workers
must have been launched with ``--disaggregation-mode`` and
``--disaggregation-bootstrap-port`` for the PD role to be
picked up (see ``model_pool.spawn_worker``); modern SGLang is
assumed.
"""
self._launch(
self._build_config(
model_id=model_id,
tokenizer_path=tokenizer_path,
urls=list(prefill_urls) + list(decode_urls),
policy=policy,
extra_models=[],
),
timeout=timeout,
)
# ----- shutdown --------------------------------------------------------
def shutdown(self) -> None:
"""SIGTERM the router; SIGKILL after 30s. Idempotent."""
if self.process is not None and self.process.poll() is None:
try:
self.process.send_signal(signal.SIGTERM)
try:
self.process.wait(timeout=30)
except subprocess.TimeoutExpired:
self.process.kill()
self.process.wait()
except ProcessLookupError:
pass
self.process = None
if self._config_path and self._config_path.exists():
self._config_path.unlink(missing_ok=True)
self._config_path = None
self._started = False
# Tear down any owned upstream workers.
for w in self._owned_workers:
if w.poll() is None:
try:
w.send_signal(signal.SIGTERM)
try:
w.wait(timeout=30)
except subprocess.TimeoutExpired:
w.kill()
w.wait()
except ProcessLookupError:
pass
self._owned_workers.clear()
# ----- HTTP introspection helpers -------------------------------------
def healthy(self, timeout: float = 5.0) -> bool:
try:
resp = httpx.get(f"{self.base_url}/healthz", timeout=timeout)
return resp.status_code == 200
except (httpx.RequestError, httpx.TimeoutException):
return False
def ready(self, timeout: float = 5.0) -> bool:
try:
resp = httpx.get(f"{self.base_url}/readyz", timeout=timeout)
return resp.status_code == 200
except (httpx.RequestError, httpx.TimeoutException):
return False
def metrics_text(self, timeout: float = 5.0) -> str | None:
try:
resp = httpx.get(f"{self.base_url}/metrics", timeout=timeout)
if resp.status_code == 200:
return resp.text
return None
except (httpx.RequestError, httpx.TimeoutException):
return None
# ----- internals ------------------------------------------------------
def _build_config(
self,
*,
model_id: str,
tokenizer_path: str,
urls: list[str],
policy: str,
extra_models: list[dict],
) -> str:
resolved_tokenizer = _resolve_tokenizer_path(tokenizer_path)
extra_model_toml = ""
for em in extra_models:
extra_model_toml += (
f'\n[[models]]\nid = "{em["id"]}"\n'
f'tokenizer_path = "{_resolve_tokenizer_path(em["tokenizer_path"])}"\n'
f'policy = "{em.get("policy", policy)}"\n'
)
# Optional tunables — only emit the [proxy] and [active_load]
# sections if a test has overridden them, so production defaults
# apply otherwise.
proxy_section = ""
if self.proxy_request_timeout_secs is not None:
proxy_section = (
f"\n[proxy]\nrequest_timeout_secs = {self.proxy_request_timeout_secs}\n"
)
active_load_section = ""
if self.stale_request_timeout_secs is not None:
active_load_section = (
f"\n[active_load]\nstale_request_timeout_secs = "
f"{self.stale_request_timeout_secs}\n"
)
urls_toml = ", ".join(f'"{u}"' for u in urls)
return f"""\
[server]
host = "{self.host}"
port = {self.port}
[[models]]
id = "{model_id}"
tokenizer_path = "{resolved_tokenizer}"
policy = "{policy}"
{extra_model_toml}
[discovery]
backend = "static_urls"
[discovery.static_urls]
urls = [{urls_toml}]
{proxy_section}{active_load_section}"""
def _launch(self, config_text: str, *, timeout: float) -> None:
if not self.binary.exists():
raise RuntimeError(
f"sgl-router binary not found at {self.binary}. "
"Build it first: `cd experimental/sgl-router && cargo build --release` "
"or set SGL_ROUTER_BINARY to the binary path."
)
# Write the main config.
fd, path = tempfile.mkstemp(suffix=".toml", prefix="sgl-router-")
os.close(fd)
self._config_path = Path(path)
self._config_path.write_text(config_text, encoding="utf-8")
logger.info("sgl-router config: %s", self._config_path)
logger.debug("sgl-router config text:\n%s", config_text)
self.process = subprocess.Popen(
[str(self.binary), "--config", str(self._config_path)],
stdout=subprocess.PIPE,
stderr=subprocess.STDOUT,
start_new_session=True,
)
try:
self._wait_ready(timeout=timeout)
except Exception:
self.shutdown()
raise
self._started = True
def _wait_ready(self, *, timeout: float) -> None:
deadline = time.time() + timeout
last_exc: Exception | None = None
while time.time() < deadline:
if self.process is not None and self.process.poll() is not None:
# Process exited early — surface stdout/stderr.
out = b""
try:
if self.process.stdout is not None:
out = self.process.stdout.read() or b""
except Exception: # noqa: BLE001
pass
raise RuntimeError(
f"sgl-router exited during startup with code "
f"{self.process.returncode}. output:\n{out.decode(errors='replace')}",
)
try:
resp = httpx.get(f"{self.base_url}/readyz", timeout=2.0)
if resp.status_code == 200:
return
except (httpx.RequestError, httpx.TimeoutException) as exc:
last_exc = exc
time.sleep(0.5)
raise TimeoutError(
f"sgl-router did not become ready at {self.base_url} within {timeout}s "
f"(last error: {last_exc})"
)
@@ -0,0 +1,228 @@
"""Minimal SGLang worker spawner for sgl-router e2e tests.
Adapted from SMG's e2e_test/infra/model_pool.py — the 1200-line original
manages a pool of long-lived workers across many tests; here we only
need a thin wrapper around ``sglang.launch_server`` that:
- allocates GPU(s) for the worker (via ``CUDA_VISIBLE_DEVICES``),
- spawns ``python3 -m sglang.launch_server`` with the right args,
- waits for ``/health`` to come up,
- optionally injects ``--kv-events-config`` so the worker exposes
the ``kv_events`` block on ``/server_info``.
A test owns a ``ModelInstance`` for its duration; teardown shuts the
worker down. No cross-test pooling — the acceptance tests are slow
enough already (model load dominates) that pooling complexity wasn't
worth porting.
"""
from __future__ import annotations
import json
import logging
import os
import signal
import socket
import subprocess
import time
from dataclasses import dataclass, field
from pathlib import Path
import httpx
from .model_specs import get_model_spec
logger = logging.getLogger(__name__)
# Passthrough Jinja chat template that emits ONLY `messages[*].content`
# joined with `\n` — matching the router's cache_aware_zmq prompt
# extraction. A worker launched with
# ``--chat-template <PASSTHROUGH_CHAT_TEMPLATE_PATH>`` tokenizes the
# raw content string, so its KV-block hashes align with what the
# router computes from the same chat-completions request. Test-only.
PASSTHROUGH_CHAT_TEMPLATE_PATH = str(
Path(__file__).parent / "passthrough_chat_template.jinja"
)
def _get_open_port() -> int:
"""Allocate an ephemeral TCP port in the range [20000, 55535].
SGLang derives its internal gRPC port as ``http_port + 10000``; if the
kernel hands us an ephemeral port above 55535, that derivation overflows
65535 and ``ServerArgs.__post_init__`` rejects it. Retrying a bounded
number of times keeps us safely below the ceiling without hand-rolling
a port registry.
"""
for _ in range(50):
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
s.bind(("127.0.0.1", 0))
port = s.getsockname()[1]
if 20000 <= port <= 55535:
return port
raise RuntimeError(
"could not allocate an ephemeral port in [20000, 55535] after 50 tries; "
"SGLang derives its internal gRPC port as http_port + 10000 and "
"rejects values above 65535"
)
@dataclass
class ModelInstance:
"""A running ``sglang.launch_server`` process.
Use as a context manager:
with spawn_worker("qwen3-0.6b", gpu_ids=[0]) as inst:
httpx.post(f"{inst.url}/generate", ...)
"""
url: str
port: int
process: subprocess.Popen
model_id: str
gpu_ids: list[int] = field(default_factory=list)
kv_events_endpoint: str | None = None
def __enter__(self) -> "ModelInstance":
return self
def __exit__(self, *exc) -> None:
self.shutdown()
def shutdown(self) -> None:
if self.process is not None and self.process.poll() is None:
try:
self.process.send_signal(signal.SIGTERM)
try:
self.process.wait(timeout=60)
except subprocess.TimeoutExpired:
self.process.kill()
self.process.wait()
except ProcessLookupError:
pass
def spawn_worker(
model_id: str,
*,
gpu_ids: list[int],
port: int | None = None,
enable_kv_events: bool = False,
kv_events_port: int | None = None,
disagg_mode: str | None = None,
bootstrap_port: int | None = None,
extra_args: list[str] | None = None,
timeout: float = 600.0,
) -> ModelInstance:
"""Spawn a single ``sglang.launch_server`` and wait for ``/health``.
Args:
model_id: Key into :data:`model_specs.MODEL_SPECS`.
gpu_ids: Concrete GPU indices to bind via ``CUDA_VISIBLE_DEVICES``.
port: HTTP port; auto-assigned if None.
enable_kv_events: If True, inject ``--kv-events-config`` with a
ZMQ publisher so the router's introspection picks up the
kv_events block from ``/server_info`` (Patch 1).
kv_events_port: ZMQ publisher port. Auto-assigned if None and
``enable_kv_events`` is True.
disagg_mode: "prefill" or "decode" for PD-disagg launches; passed
through as ``--disaggregation-mode``.
bootstrap_port: PD-disagg bootstrap port (prefill side only).
extra_args: Additional CLI args appended verbatim.
timeout: Health-check timeout. Cold-start on a fresh GPU can be
slow; default is 10 minutes.
"""
spec = get_model_spec(model_id)
port = port or _get_open_port()
base_url = f"http://127.0.0.1:{port}"
cmd = [
"python3",
"-m",
"sglang.launch_server",
"--model-path",
spec["model"],
"--port",
str(port),
"--host",
"127.0.0.1",
"--tp",
str(spec.get("tp", 1)),
]
cmd.extend(spec.get("worker_args", []) or [])
kv_events_endpoint: str | None = None
if enable_kv_events:
kv_port = kv_events_port or _get_open_port()
kv_events_endpoint = f"tcp://*:{kv_port}"
kv_cfg = {
"publisher": "zmq",
"endpoint": kv_events_endpoint,
"topic": "kv",
}
cmd.extend(["--kv-events-config", json.dumps(kv_cfg)])
if disagg_mode is not None:
cmd.extend(["--disaggregation-mode", disagg_mode])
if bootstrap_port is not None:
cmd.extend(["--disaggregation-bootstrap-port", str(bootstrap_port)])
if extra_args:
cmd.extend(extra_args)
env = os.environ.copy()
env["CUDA_VISIBLE_DEVICES"] = ",".join(str(g) for g in gpu_ids)
logger.info(
"spawning sglang worker: model=%s port=%d gpus=%s disagg=%s",
model_id,
port,
gpu_ids,
disagg_mode,
)
proc = subprocess.Popen(
cmd,
env=env,
stdout=subprocess.PIPE,
stderr=subprocess.STDOUT,
start_new_session=True,
)
inst = ModelInstance(
url=base_url,
port=port,
process=proc,
model_id=model_id,
gpu_ids=list(gpu_ids),
kv_events_endpoint=kv_events_endpoint,
)
# Wait for /health. Cold-start on H200 with weights uncached can take
# ~5 minutes; CI configurations should pre-warm.
deadline = time.time() + timeout
while time.time() < deadline:
if proc.poll() is not None:
out = b""
try:
if proc.stdout is not None:
out = proc.stdout.read() or b""
except Exception: # noqa: BLE001
pass
raise RuntimeError(
f"sglang worker exited during startup with code {proc.returncode}; "
f"cmd: {' '.join(cmd)}\noutput:\n{out.decode(errors='replace')}",
)
try:
resp = httpx.get(f"{base_url}/health", timeout=2.0)
if resp.status_code == 200:
logger.info("sglang worker ready at %s", base_url)
return inst
except (httpx.RequestError, httpx.TimeoutException):
pass
time.sleep(2.0)
inst.shutdown()
raise TimeoutError(
f"sglang worker did not become healthy at {base_url} within {timeout}s",
)
@@ -0,0 +1,77 @@
"""Model specifications for sgl-router e2e tests.
Adapted from SMG's e2e_test/infra/model_specs.py. The same dict-of-dicts
shape (so test code reads the same) but the entries are narrower —
sgl-router tests today target small/medium models only; the larger
function-calling / reasoning models from SMG are out of scope.
Each entry:
- model: HuggingFace path or local path (env-resolved)
- memory_gb: estimated single-GPU footprint
- tp: tensor-parallel size (= GPUs needed)
- features: feature tags for filtering
- worker_args: optional extra `sglang.launch_server` flags
"""
from __future__ import annotations
import os
# Local-cache root for CI / cluster nodes that pre-download HF weights.
# Mirrors the SMG `ROUTER_LOCAL_MODEL_PATH` env var.
ROUTER_LOCAL_MODEL_PATH = os.environ.get("ROUTER_LOCAL_MODEL_PATH", "")
def _resolve_model_path(hf_path: str) -> str:
"""Prefer a local copy of the model when one exists under
``ROUTER_LOCAL_MODEL_PATH``; otherwise fall back to the HuggingFace ID.
"""
if ROUTER_LOCAL_MODEL_PATH:
local_path = os.path.join(ROUTER_LOCAL_MODEL_PATH, hf_path)
if os.path.exists(local_path):
return local_path
return hf_path
MODEL_SPECS: dict[str, dict] = {
# Fast-start tiny model for convergence / decode-affinity / stale-request
# tests. Single GPU, ~2 GB weights, sub-30s start on a warm cache.
"qwen3-0.6b": {
"model": _resolve_model_path("Qwen/Qwen3-0.6B"),
"memory_gb": 4,
"tp": 1,
"features": ["chat", "streaming"],
},
# Standard small chat model — matches SMG's `llama-1b` entry.
"llama-1b": {
"model": _resolve_model_path("meta-llama/Llama-3.2-1B-Instruct"),
"memory_gb": 4,
"tp": 1,
"features": ["chat", "streaming"],
},
# Primary 8B chat model — matches SMG's `llama-8b`.
"llama-8b": {
"model": _resolve_model_path("meta-llama/Llama-3.1-8B-Instruct"),
"memory_gb": 16,
"tp": 1,
"features": ["chat", "streaming"],
},
}
def get_model_spec(model_id: str) -> dict:
"""Return the spec dict for ``model_id``; KeyError if absent."""
if model_id not in MODEL_SPECS:
raise KeyError(
f"Unknown model: {model_id}. Available: {list(MODEL_SPECS.keys())}"
)
return MODEL_SPECS[model_id]
def get_models_with_feature(feature: str) -> list[str]:
"""Filter model IDs by feature tag (e.g. ``streaming``, ``chat``)."""
return [
model_id
for model_id, spec in MODEL_SPECS.items()
if feature in spec.get("features", [])
]
@@ -0,0 +1,13 @@
{#-
Passthrough chat template for cache-aware-zmq e2e tests.
Emits ONLY `messages[*].content` joined with `\n` — no role markers,
no special tokens, no generation prompt. This is the SAME shape the
router's cache_aware_zmq policy produces in `extract_prompt_text`,
so a worker launched with `--chat-template <this file>` tokenizes the
same string the router will tokenize for routing — making block
hashes align across worker KV cache and router HashTree.
Use only for tests; not appropriate for any real chat workload.
-#}
{{- messages | map(attribute='content') | join('\n') -}}