From 97ba99067dc1915771edc96972a8dbd08dc21148 Mon Sep 17 00:00:00 2001 From: Shangming Cai Date: Thu, 27 Aug 2026 19:33:18 +0800 Subject: [PATCH] Publish per-scheduler load on a dedicated socket for load-aware routers (#34608) Co-authored-by: Kangyan Zhou Co-authored-by: Zhangheng --- .../advanced_features/server_arguments.mdx | 8 +- python/sglang/srt/disaggregation/kv_events.py | 173 +++++- python/sglang/srt/managers/scheduler.py | 69 ++- .../kv_events_publisher.py | 8 +- .../scheduler_components/load_publisher.py | 267 ++++++++++ python/sglang/srt/server_args.py | 118 +++- .../unit/disaggregation/test_kv_events.py | 84 +++ .../unit/entrypoints/test_server_info.py | 197 +++++++ .../unit/managers/test_loadstat_wire.py | 502 ++++++++++++++++++ .../managers/test_scheduler_on_idle_load.py | 52 ++ 10 files changed, 1443 insertions(+), 35 deletions(-) create mode 100644 python/sglang/srt/managers/scheduler_components/load_publisher.py create mode 100644 test/registered/unit/managers/test_loadstat_wire.py create mode 100644 test/registered/unit/managers/test_scheduler_on_idle_load.py diff --git a/docs/docs/advanced_features/server_arguments.mdx b/docs/docs/advanced_features/server_arguments.mdx index 9ecfff75a..496000c89 100644 --- a/docs/docs/advanced_features/server_arguments.mdx +++ b/docs/docs/advanced_features/server_arguments.mdx @@ -976,7 +976,13 @@ Please consult the documentation below and [server_args.py](https://github.com/s `--kv-events-config` - Config in json format for NVIDIA dynamo KV event publishing. Publishing will be enabled if this flag is used. + Config in json format for NVIDIA dynamo KV event publishing. Publishing will be enabled if this flag is used. Runtime-load publishing for load-aware routers is a separate opt-in; see `--load-publish-endpoint`. + None + Type: str + + + `--load-publish-endpoint` + Opt in to the runtime-load PUB socket that load-aware routers subscribe to. Off by default (unset or `off`). Use `auto` to reserve the dp_size ports packed after the `--kv-events-config` range, or a wildcard-host TCP address (e.g. `tcp://*:6000`) to place it explicitly; rank r binds port+r and `/server_info` advertises the base under the `kv_events` block. Requires `--kv-events-config`; startup fails if set without one, not bindable, or overlapping the KV range. `auto` reserves 2*dp_size ports from the KV base — space co-hosted engines accordingly. The router-facing update cadence follows `--load-snapshot-publish-interval`, so a large value there also staleness-caps this feed. None Type: str diff --git a/python/sglang/srt/disaggregation/kv_events.py b/python/sglang/srt/disaggregation/kv_events.py index 5113d4759..b03b1efd8 100644 --- a/python/sglang/srt/disaggregation/kv_events.py +++ b/python/sglang/srt/disaggregation/kv_events.py @@ -27,12 +27,17 @@ from abc import ABC, abstractmethod from collections import deque from itertools import count from queue import Queue -from typing import Any, Callable, Optional, Union +from typing import TYPE_CHECKING, Any, Callable, Optional, Union import msgspec import zmq from pydantic import BaseModel +from sglang.srt.utils.network import NetworkAddress + +if TYPE_CHECKING: + from sglang.srt.distributed.parallel_state_wrapper import ParallelState + logger = logging.getLogger(__name__) @@ -58,6 +63,172 @@ def select_kv_publisher_dp_rank( return dp_rank or 0 +def is_kv_publisher_rank(kv_events_config: Optional[str], ps: "ParallelState") -> bool: + """Whether this scheduler owns a KV-event publisher slot: one per + independent KV cache (pp/attn-TP/attn-CP rank 0). Shared by + `SchedulerKvEventsPublisher` and `SchedulerLoadPublisher`, which must + gate identically or their /server_info-derived ports disagree. + """ + return bool( + kv_events_config + and ps.pp_rank == 0 + and ps.attn_tp_rank == 0 + and ps.attn_cp_rank == 0 + ) + + +# Advertised as `load_topic` in /server_info; the load socket carries only +# load, so subscribers can subscribe-all. +LOAD_TOPIC = "load" + +# Hosts a PUB socket binds rather than connects to. Matched on the parsed +# host, not a substring: "::" appears inside every IPv6 address, so a +# substring test would wrongly call a concrete remote host bindable. +_BIND_WILDCARD_HOSTS = frozenset({"*", "0.0.0.0", "::"}) + + +def parse_tcp_port(endpoint: Optional[str]) -> Optional[int]: + """Legal port of a tcp:// endpoint regardless of host, or None. + + Host-agnostic: answers "which ports does something else occupy" for the + collision checks (the replay ROUTER binds any host spelling). + """ + if not endpoint or not endpoint.startswith("tcp://"): + return None + try: + port = NetworkAddress.parse(endpoint[len("tcp://") :]).port + except ValueError: + return None + return port if 0 < port <= 65535 else None + + +def parse_advertisable_tcp(endpoint: Optional[str]) -> Optional[tuple[str, int]]: + """``(host, port)`` of a tcp:// endpoint fit for /server_info, else None. + + Any host (KV events work connect-style); IPv6 re-bracketed so consumers + can splice ``tcp://{host}:{port}``. Bare unbracketed IPv6 is rejected — + same parse as the resolver, so descriptor and bind agree. + """ + if not endpoint or not endpoint.startswith("tcp://"): + return None + try: + addr = NetworkAddress.parse(endpoint[len("tcp://") :]) + except ValueError: + return None + if not addr.host or not (0 < addr.port <= 65535): + return None + host = f"[{addr.host}]" if addr.is_ipv6 else addr.host + return host, addr.port + + +def parse_bindable_tcp(endpoint: Optional[str]) -> Optional[tuple[str, int]]: + """``(host, port)`` if a PUB socket can BIND this tcp:// endpoint, else + None. A concrete host is connect-style here, so a load PUB there would + reach nobody while reporting no error.""" + if not endpoint or not endpoint.startswith("tcp://"): + return None + try: + addr = NetworkAddress.parse(endpoint[len("tcp://") :]) + except ValueError: + return None + if addr.host not in _BIND_WILDCARD_HOSTS or not (0 < addr.port <= 65535): + return None + return addr.host, addr.port + + +def resolve_load_pub_range( + *, + kv_endpoint: Optional[str], + replay_endpoint: Optional[str], + dp_size: int, + load_publish_endpoint: Optional[str] = None, +) -> tuple[Optional[tuple[str, int]], Optional[str]]: + """``((host, base), reason)`` for the load PUB range — exactly one is None. + + Rank ``r`` binds ``base + r`` and ``/server_info`` advertises ``base``. + Single source of truth for both the bind (`SchedulerLoadPublisher`) and + the advertisement (`describe_kv_events_publisher`), so they cannot drift. + + Opt-in via ``--load-publish-endpoint``: unset (or ``off``) disables it, so + an upgrade never reserves a port a co-hosted neighbor's KV publisher would + bind. ``auto`` packs the range after the KV-event range, bumping past an + overlapping replay ROUTER range (with the conventional replay = kv + 1, + always); an explicit ``tcp://`` address sets it outright. + + ``reason`` is set when an operator would want to know why publishing is + off (unusable endpoint, collision, u16 overflow) and None when the decline + is unremarkable (feature off). Callers log it once; /server_info calls + this per request, so it must not log here. + + Two inherited limits, both from the KV-event discovery structure: with + ``page_size`` <= 0 `describe_kv_events_publisher` suppresses the whole + block, so the range binds unadvertised; and with DP-attention across + ``nnodes`` > 1 the single advertised base is paired with one worker-URL + host, so ranks on other nodes are unreachable at that host. + """ + # Opt-in: off unless the operator sets `auto` (derive) or an address, so an + # upgrade never claims a port a co-hosted neighbor's KV publisher binds. + mode = (load_publish_endpoint or "").strip() + if dp_size < 1 or not mode or mode.lower() == "off": + return None, None + + if mode.lower() == "auto": + resolved = parse_bindable_tcp(kv_endpoint) + if resolved is None: + why = ( + "--kv-events-config is not set" + if kv_endpoint is None + else f"{kv_endpoint!r} is not one" + ) + return None, ( + f"--load-publish-endpoint=auto needs a bindable wildcard-host " + f"tcp:// --kv-events-config endpoint to pack after; {why}" + ) + host, kv_base = resolved + base = kv_base + dp_size + replay_base = parse_tcp_port(replay_endpoint) + if ( + replay_base is not None + and base < replay_base + dp_size + and replay_base < base + dp_size + ): + # Overlap implies kv < replay < kv + 2*dp_size, so packing after + # the replay range also clears the KV range. + base = replay_base + dp_size + else: + # Explicit address. Discovery still needs the kv_events block, absent + # for a non-tcp KV endpoint — so the range would bind but never + # advertise. + if parse_tcp_port(kv_endpoint) is None: + absent = ( + "without --kv-events-config" + if kv_endpoint is None + else f"for endpoint {kv_endpoint!r}" + ) + return None, ( + f"--load-publish-endpoint={mode!r} needs a routable tcp:// " + f"--kv-events-config endpoint: routers discover the load range " + f"through /server_info's kv_events block, absent {absent}, so " + f"the socket would be bound but never advertised" + ) + resolved = parse_bindable_tcp(mode) + if resolved is None: + return None, ( + f"--load-publish-endpoint={mode!r} is not a bindable tcp:// " + f"address (a concrete host would be connected to, not bound)" + ) + host, base = resolved + for port in (parse_tcp_port(kv_endpoint), parse_tcp_port(replay_endpoint)): + if port is not None and base < port + dp_size and port < base + dp_size: + return None, ( + f"--load-publish-endpoint range [{base}, {base + dp_size}) " + f"overlaps the kv-events range [{port}, {port + dp_size})" + ) + if base + dp_size - 1 > 65535: + return None, f"load port range from {base} would run past the u16 ceiling" + return (host, base), None + + class EventBatch( msgspec.Struct, array_like=True, # type: ignore[call-arg] diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index dfe597381..7254ddca3 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -228,6 +228,9 @@ from sglang.srt.managers.scheduler_components.kv_events_publisher import ( SchedulerKvEventsPublisher, ) from sglang.srt.managers.scheduler_components.load_inquirer import SchedulerLoadInquirer +from sglang.srt.managers.scheduler_components.load_publisher import ( + SchedulerLoadPublisher, +) from sglang.srt.managers.scheduler_components.logprob_result_processor import ( SchedulerLogprobResultProcessor, ) @@ -361,6 +364,11 @@ TEST_RETRACT_NO_PREFILL_BS = envs.SGLANG_TEST_RETRACT_NO_PREFILL_BS.get() STEP_MAX_US = 2_000_000 +# Min wall-clock between load publishes on the stalled no-batch path, which +# spins on_idle without sleeping. Bounds the O(queue) get_loads for both the +# DP-balancing writer and the router-facing socket. +LOAD_STALL_REFRESH_S = 0.05 + def _accumulate_decode_moment( totals: list[float], @@ -395,6 +403,10 @@ class Scheduler( ): """A scheduler that manages a tensor parallel GPU worker.""" + # Class-level default so on_idle's stall gate works even if a fork + # overrides init_load_publisher (which would otherwise not set it). + _last_stall_publish_ts: float = float("-inf") + def __init__( self, server_args: ServerArgs, @@ -658,6 +670,8 @@ class Scheduler( self.init_kv_events_publisher() + self.init_load_publisher() + self.init_load_inquirer() self.init_output_streamer() @@ -802,18 +816,24 @@ class Scheduler( self.idle_sleeper = None def publish_load_snapshot(self, force: bool = False): + """Returns the LoadSnapshot it published, or None when disabled, + throttled, or failed — so co-located sinks (the router-facing load + publisher) can reuse it instead of walking the queues again.""" writer = self.load_snapshot_writer if writer is None: - return + return None if not force: writer.publish_counter += 1 if writer.publish_counter < writer.publish_interval: - return + return None writer.publish_counter = 0 try: - writer.write(self.load_inquirer.get_loads()) + load = self.load_inquirer.get_loads() + writer.write(load) + return load except Exception as e: logger.warning("load snapshot publish failed: %s", e) + return None def init_tokenizer(self): server_args = self.server_args @@ -2158,6 +2178,18 @@ class Scheduler( get_stats=lambda: self.metrics_reporter.stats, ) + def init_load_publisher(self) -> None: + # Router-facing load reporting; rank gating and no-op fallback live + # inside the component. Same interval as the DP-balancing writer so + # the two fire in phase and the load sink always reuses that snapshot + # instead of walking the queues itself. + self.load_publisher = SchedulerLoadPublisher( + kv_events_config=get_observability().kv_events_config, + ps=self.ps, + load_publish_endpoint=get_observability().load_publish_endpoint, + publish_interval=get_observability().load_snapshot_publish_interval, + ) + def init_load_inquirer(self) -> None: self.total_prefill_uncached_tokens = 0 self.total_prefill_busy_us = 0 @@ -4119,7 +4151,14 @@ class Scheduler( # Flush async trace ops here: in overlap mode this CPU work runs while # the next batch's GPU forward is in flight, giving free overlap. flush_trace_batch(batch.reqs) - self.publish_load_snapshot(force=batch.forward_mode.is_extend()) + snapshot = self.publish_load_snapshot(force=batch.forward_mode.is_extend()) + # Router-facing gauge on the dedicated PUB socket, reusing the + # snapshot above rather than walking the queues again. + self.load_publisher.publish_load_stat( + self.load_inquirer.get_loads, + force=batch.forward_mode.is_extend(), + snapshot=snapshot, + ) if batch.forward_mode.is_decode(): self.batch_result_processor.process_batch_result_decode(batch, result) @@ -4235,7 +4274,21 @@ class Scheduler( # Flush any health-check signal deferred while the engine was busy. self.maybe_send_health_check_signal() + # Publish before the fully-idle gate: a no-batch-but-not-idle stall + # (queues parked under KV pressure / disagg transfer) has no + # process_batch_result to publish the growing gauge, and gating here + # froze /get_loads, DP balancing, and the LoadStat for the stall. This + # path spins without sleeping, so a wall-clock floor bounds the + # O(queue) get_loads for both sinks; the fully-idle publish runs + # post-flush below. if not self.is_fully_idle(): + now = time.monotonic() + if now - self._last_stall_publish_ts >= LOAD_STALL_REFRESH_S: + self._last_stall_publish_ts = now + snapshot = self.publish_load_snapshot(force=True) + self.load_publisher.publish_load_stat( + self.load_inquirer.get_loads, force=True, snapshot=snapshot + ) return if self.enable_unified_memory: @@ -4273,8 +4326,12 @@ class Scheduler( # reset token ratio self.new_token_ratio_tracker.reset() - # Publish the idle state so /get_loads and DP balancing do not see stale load. - self.publish_load_snapshot(force=True) + # Fully-idle publish, post-flush so the gauge reflects compacted KV. + # Forced (immediate) so the busy->idle transition is never delayed. + snapshot = self.publish_load_snapshot(force=True) + self.load_publisher.publish_load_stat( + self.load_inquirer.get_loads, force=True, snapshot=snapshot + ) # sleep until next event self.maybe_sleep_on_idle() diff --git a/python/sglang/srt/managers/scheduler_components/kv_events_publisher.py b/python/sglang/srt/managers/scheduler_components/kv_events_publisher.py index 5aee67f7b..b483a9ac9 100644 --- a/python/sglang/srt/managers/scheduler_components/kv_events_publisher.py +++ b/python/sglang/srt/managers/scheduler_components/kv_events_publisher.py @@ -15,6 +15,7 @@ import zmq from sglang.srt.disaggregation.kv_events import ( EventPublisherFactory, KVEventBatch, + is_kv_publisher_rank, select_kv_publisher_dp_rank, ) from sglang.srt.managers.io_struct import hook_custom_types, sock_send @@ -61,12 +62,7 @@ class SchedulerKvEventsPublisher: self.init_kv_events(self.kv_events_config) def init_kv_events(self, kv_events_config: Optional[str]): - self.enable_kv_cache_events = bool( - kv_events_config - and self.ps.pp_rank == 0 - and self.ps.attn_tp_rank == 0 - and self.ps.attn_cp_rank == 0 - ) + self.enable_kv_cache_events = is_kv_publisher_rank(kv_events_config, self.ps) if self.enable_kv_cache_events: self.kv_event_publisher = EventPublisherFactory.create( diff --git a/python/sglang/srt/managers/scheduler_components/load_publisher.py b/python/sglang/srt/managers/scheduler_components/load_publisher.py new file mode 100644 index 000000000..90a28b73d --- /dev/null +++ b/python/sglang/srt/managers/scheduler_components/load_publisher.py @@ -0,0 +1,267 @@ +"""Per-scheduler load reporting for load-aware routers. + +Each scheduler publishes a periodic `LoadStat` gauge on its own ZMQ PUB +socket so out-of-process load-aware routers (e.g. sgl-router's +`cache_aware_zmq` policy) can route on real queue depth instead of a +router-side in-flight counter. The in-deployment counterpart lives in +`sglang.srt.managers.load_snapshot` (SHM / PUSH to node 0), which a router +that only knows the worker URL cannot subscribe to; the port is instead +advertised via `/server_info` (`ServerArgs.describe_kv_events_publisher`). +The payload is a compact tagged subset of `LoadSnapshot` so the wire +contract stays fixed as the snapshot grows. + +Framing is the KV-event socket's, so one subscriber loop handles both: +``[b"load", big-endian i64 seq, msgpack LoadStat]``. The transport is a +plain synchronous PUB socket (a send just enqueues to ZMQ's IO thread) — +no background thread or replay buffer, which a gauge does not need. + +Opt-in via `--load-publish-endpoint` (`auto` to pack after the KV range, or +an explicit address); off by default so an upgrade never reserves a port a +co-hosted neighbor's KV publisher binds. The port comes from +`resolve_load_pub_range` (the same function `/server_info` advertises with, +so the two cannot drift). With `auto`, a worker's ZMQ footprint is +`2 * dp_size` ports after its KV base (`2 * dp_size + 1` with the +conventional adjacent replay), so co-hosted workers must space their KV +bases that far apart or move the range with an explicit address. +""" + +from __future__ import annotations + +import atexit +import logging +import time +from itertools import count +from typing import TYPE_CHECKING, Callable, Optional + +import msgspec +import zmq + +from sglang.srt.disaggregation.kv_events import ( + LOAD_TOPIC, + KVEventsConfig, + is_kv_publisher_rank, + resolve_load_pub_range, + select_kv_publisher_dp_rank, +) +from sglang.srt.utils.network import NetworkAddress, is_zmq_endpoint_ipv6 + +if TYPE_CHECKING: + from sglang.srt.distributed.parallel_state_wrapper import ParallelState + from sglang.srt.managers.load_snapshot import LoadSnapshot + +logger = logging.getLogger(__name__) + +# Default call throttle (overridden with the DP-snapshot interval so the two +# sinks fire in phase). Publish at most once per this many calls unless force. +LOAD_PUBLISH_INTERVAL = 5 + +# An unchanged stat is re-sent at most this often; a changed one always goes +# out immediately, so transitions are never delayed. Bounds the send rate on +# the idle spin loop (on_idle force-publishes every iteration). +LOAD_PUBLISH_HEARTBEAT_S = 1.0 + +# Small HWM: load is a gauge, so shedding at a full pipe loses readings the +# next heartbeat supersedes. ZMQ_CONFLATE (true newest-wins) is unusable — it +# keeps a single frame, breaking the 3-frame framing — so a bounded backlog +# is the closest fit. +LOAD_PUB_HWM = 8 + +_encoder = msgspec.msgpack.Encoder() + + +class LoadStat( + msgspec.Struct, + array_like=True, # type: ignore[call-arg] + # No omit_defaults: it may trim trailing defaults, shortening a shape the + # router decodes positionally. + gc=False, # type: ignore[call-arg] + tag=True, # type: ignore[call-arg] +): + """Per-scheduler runtime load snapshot. + + Wire shape (tag + array_like): ``["LoadStat", num_running_reqs, + num_waiting_reqs, num_tokens, max_total_num_tokens, attn_dp_rank]``. The + router reads the four counts; array_like always emits the trailing field + (null when unset), so a decoder must tolerate it. + """ + + num_running_reqs: int + num_waiting_reqs: int + num_tokens: int # KV tokens in use + max_total_num_tokens: int # KV capacity; 0 when unknown + # attn_dp_rank under DP attention, else the plain dp_rank; informational + # only (the router keys by socket rank). Name follows EventBatch's. + attn_dp_rank: Optional[int] = None + + +def _open_pub_socket(endpoint: str) -> zmq.Socket: + """Bind the load PUB socket. Module-level so tests can stub the one side + effect while exercising the real gating and port derivation. Not + get_zmq_socket: that sets SNDHWM=0, defeating LOAD_PUB_HWM.""" + sock = zmq.Context.instance().socket(zmq.PUB) + try: + sock.set_hwm(LOAD_PUB_HWM) + sock.setsockopt(zmq.LINGER, 0) + if is_zmq_endpoint_ipv6(endpoint): + sock.setsockopt(zmq.IPV6, 1) + sock.bind(endpoint) + except Exception: + sock.close() # don't leak the handle on the shared context + raise + return sock + + +class SchedulerLoadPublisher: + """Owns one scheduler's dedicated load PUB socket and the throttled, + best-effort `publish_load_stat` path. + + Enabled on the same condition as KV-event publishing + (`is_kv_publisher_rank`), keyed per rank like it + (`select_kv_publisher_dp_rank`) so pure-DP replicas don't collide. Stays + a no-op (`_socket is None`) when disabled or no range is resolvable. + """ + + def __init__( + self, + *, + kv_events_config: Optional[str], + ps: ParallelState, + load_publish_endpoint: Optional[str] = None, + publish_interval: int = LOAD_PUBLISH_INTERVAL, + ) -> None: + # _socket is None == disabled: every early return below leaves it so, + # and publish_load_stat then skips the snapshot entirely. + self._socket: Optional[zmq.Socket] = None + self._rank = 0 + self._seq = count() + self._publish_counter = 0 + self._publish_interval = max(1, publish_interval) + # Last sent counts + timestamp, driving the dedup/heartbeat. + self._last_counts: Optional[tuple] = None + self._last_publish_ts = 0.0 + self._publish_failed = False + if not is_kv_publisher_rank(kv_events_config, ps): + return + try: + cfg = KVEventsConfig.from_cli(kv_events_config) + except Exception: + # Malformed config: the KV publisher would fail too; stay a no-op. + return + if cfg.publisher == "null" or not cfg.endpoint: + # KV publishing off. Silent: an opted-in server already failed in + # check_load_publish_args, so reaching here means load publishing + # wasn't requested (this runs before the resolver sees the mode). + return + # Same resolver /server_info advertises with, so a router never + # subscribes to a range this declines — except a runtime bind failure + # below, which the advertisement can't retract (router sees silence). + resolved, reason = resolve_load_pub_range( + kv_endpoint=cfg.endpoint, + replay_endpoint=cfg.replay_endpoint, + dp_size=ps.dp_size, + load_publish_endpoint=load_publish_endpoint, + ) + if resolved is None: + if reason: + logger.warning("load-publisher disabled: %s", reason) + return + host, base = resolved + self._rank = select_kv_publisher_dp_rank( + ps.attn_dp_size, ps.attn_dp_rank, ps.dp_rank + ) + endpoint = NetworkAddress(host, base + self._rank).to_tcp() + try: + self._socket = _open_pub_socket(endpoint) + # No scheduler shutdown hook to close() from; LINGER=0 keeps a + # hard exit safe. (The KV-event publisher cleans up the same way.) + atexit.register(self.close) + except Exception: + # Best-effort: a bind failure must not take down startup. + logger.warning( + "load-publisher disabled: failed to bind the load socket at " + "%r; /server_info advertises this range but nothing is " + "listening on it", + endpoint, + exc_info=True, + ) + + @property + def enable(self) -> bool: + """True when a real load PUB socket is bound.""" + return self._socket is not None + + def publish_load_stat( + self, + load_provider: Callable[[], LoadSnapshot], + force: bool = False, + snapshot: Optional[LoadSnapshot] = None, + ) -> None: + """Publish a load snapshot, throttled to `publish_interval` calls + unless `force`; an unchanged stat is re-sent at most once per + [`LOAD_PUBLISH_HEARTBEAT_S`], a changed one always immediately. + + `load_provider` reads live scheduler state + (`SchedulerLoadInquirer.get_loads`), used over metrics stats which + are only populated under `--enable-metrics`. Skipped when the caller + passes `snapshot` (one it already computed for the DP-balancing sink + this cycle). + + Best-effort: never crashes the loop (routers fall back to their own + counter). + """ + if self._socket is None: + return + + self._publish_counter += 1 + if not force and self._publish_counter < self._publish_interval: + return + # Reset where the throttle passes, not on send: a dedup hit or + # provider failure would otherwise leave it saturated, silently + # disengaging the throttle onto the O(queue) provider every step. + self._publish_counter = 0 + + now = time.monotonic() + try: + load = snapshot if snapshot is not None else load_provider() + counts = ( + load.num_running_reqs, + load.num_waiting_reqs, + load.num_used_tokens, + load.max_total_num_tokens, + ) + if ( + counts == self._last_counts + and now - self._last_publish_ts < LOAD_PUBLISH_HEARTBEAT_S + ): + return + payload = _encoder.encode( + LoadStat( + num_running_reqs=counts[0], + num_waiting_reqs=counts[1], + num_tokens=counts[2], + max_total_num_tokens=counts[3], + attn_dp_rank=self._rank, + ) + ) + seq = next(self._seq).to_bytes(8, "big") + # PUB never blocks — it sheds at HWM. A silently dropped reading is + # superseded by the next heartbeat. + self._socket.send_multipart((LOAD_TOPIC.encode(), seq, payload)) + self._last_counts = counts + self._last_publish_ts = now + self._publish_failed = False + except Exception: + # Never crash the scheduler loop over a routing hint; log once per + # failure episode (this runs every loop, so don't flood). + if not self._publish_failed: + self._publish_failed = True + logger.warning( + "load-publisher: publish failed; routers fall back to " + "their in-flight load signal", + exc_info=True, + ) + + def close(self) -> None: + if self._socket is not None: + self._socket.close() + self._socket = None diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 8fbe685c3..71937631a 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -1592,7 +1592,12 @@ class ServerArgs: ] = False kv_events_config: A[ Optional[str], - "Config in json format for NVIDIA dynamo KV event publishing. Publishing will be enabled if this flag is used.", + "Config in json format for NVIDIA dynamo KV event publishing. Publishing will be enabled if this flag is used. Runtime-load publishing for load-aware routers is a separate opt-in; see --load-publish-endpoint.", + NS("observability"), + ] = None + load_publish_endpoint: A[ + Optional[str], + "Opt in to the runtime-load PUB socket that load-aware routers subscribe to. Off by default (unset or 'off'). Use 'auto' to reserve the dp_size ports packed after the --kv-events-config range, or a wildcard-host TCP address (e.g. tcp://*:6000) to place it explicitly; rank r binds port+r and /server_info advertises the base under the kv_events block. Requires --kv-events-config to describe a publisher (routers discover the base through /server_info); startup fails if this is set without one, is not bindable, or overlaps the KV range. Note: 'auto' reserves 2*dp_size ports from the KV base — space co-hosted engines accordingly. The router-facing update cadence follows --load-snapshot-publish-interval (shared to avoid double-collecting the snapshot), so a large value there also staleness-caps this feed.", NS("observability"), ] = None enable_forward_pass_metrics: A[ @@ -10260,6 +10265,50 @@ class ServerArgs: "--kv-canary-sweep-interval requires --kv-canary in {log, raise}" ) + self.check_load_publish_args() + + def check_load_publish_args(self): + """Fail fast at the entrypoint on a --load-publish-endpoint the + scheduler would decline (no active kv-events publisher to advertise + through, unbindable, overlapping the KV range, u16 overflow) rather + than only warning — or silently doing nothing — from a scheduler + subprocess. Routes through the same resolver the scheduler binds and + /server_info advertises with.""" + mode = (self.load_publish_endpoint or "").strip() + if not mode or mode.lower() == "off": + return # disabled; nothing to validate + + server_cfg = resolving_view(self) + + from sglang.srt.disaggregation.kv_events import ( + KVEventsConfig, + resolve_load_pub_range, + ) + + if not self.kv_events_config: + raise ValueError( + "--load-publish-endpoint requires --kv-events-config: routers" + " discover the load range through /server_info's kv_events" + " block, absent without a publisher." + ) + try: + cfg = KVEventsConfig.from_cli(self.kv_events_config) + except Exception as e: + raise ValueError(f"--kv-events-config is not parseable: {e}") + if cfg.publisher == "null" or not cfg.endpoint: + raise ValueError( + "--load-publish-endpoint needs an active --kv-events-config" + " publisher; got publisher='null' or an empty endpoint." + ) + _, reason = resolve_load_pub_range( + kv_endpoint=cfg.endpoint, + replay_endpoint=cfg.replay_endpoint, + dp_size=server_cfg.dp_size, + load_publish_endpoint=mode, + ) + if reason: + raise ValueError(reason) + def check_lora_server_args(self): cfg = resolving_view(self) @@ -10635,6 +10684,19 @@ class ServerArgs: # DCP shards within a rank # rather than adding # publishers + "load_endpoint_port_base": , + # base TCP port of the load + # range (load rank r = base + # + r). Consumers MUST read + # this key, not re-derive + # it; present only when + # --load-publish-endpoint + # opted in and a range + # resolved + "load_topic": "load", # SUB filter for the load + # socket; present iff + # load_endpoint_port_base + # is present } Returns None (i.e. "no publisher to describe") when any of: @@ -10645,17 +10707,27 @@ class ServerArgs: block_size would cause silent KV-cache misses by hashing prompts at the wrong granularity on the router side), * the endpoint is not a routable TCP address (inproc:// / - ipc://, missing port, non-integer port, or port outside - 1..65535). + ipc://, missing port, non-integer port, port outside + 1..65535, or a bare unbracketed IPv6 host, which is + ambiguous). - Reuses KVEventsConfig.from_cli for JSON parsing; the inline - rfind(":") endpoint split mirrors - ZmqEventPublisher.offset_endpoint_port rather than adding a - new module-level helper. + NOTE for load-socket consumers: pair the load port with the worker's + own URL host, as with the KV SUB endpoints — endpoint_host is a + wildcard ("*", "0.0.0.0", "::") whenever the default packing applies, + so splicing it yields tcp://*:PORT and connects to nothing. + + Reuses parse_advertisable_tcp and resolve_load_pub_range — the same + helpers the scheduler binds through — so the advertisement cannot + drift from the sockets. """ # Lazy import so loading server_args doesn't pull in # disaggregation / msgspec / zmq at module top level. - from sglang.srt.disaggregation.kv_events import KVEventsConfig + from sglang.srt.disaggregation.kv_events import ( + LOAD_TOPIC, + KVEventsConfig, + parse_advertisable_tcp, + resolve_load_pub_range, + ) resolved = resolving_view(self) raw = resolved.kv_events_config @@ -10671,21 +10743,12 @@ class ServerArgs: return None if cfg.publisher == "null" or not cfg.endpoint: return None - if not cfg.endpoint.startswith("tcp://"): - return None - body = cfg.endpoint[len("tcp://") :] - last_colon = body.rfind(":") - if last_colon < 0: - return None - host = body[:last_colon] - try: - port = int(body[last_colon + 1 :]) - except ValueError: - return None - if not host or not (0 < port < 65536): + resolved_kv = parse_advertisable_tcp(cfg.endpoint) + if resolved_kv is None: return None + host, port = resolved_kv - return { + descriptor = { "publisher": cfg.publisher, "endpoint_host": host, "endpoint_port_base": port, @@ -10693,6 +10756,19 @@ class ServerArgs: "block_size": resolved.kv_event_block_size, "dp_size": resolved.dp_size, } + # Load range, from the same resolver SchedulerLoadPublisher binds + # with (so the two can't drift). The decline reason is logged once at + # startup, not here — this runs per /server_info request. + resolved_range, _reason = resolve_load_pub_range( + kv_endpoint=cfg.endpoint, + replay_endpoint=cfg.replay_endpoint, + dp_size=resolved.dp_size, + load_publish_endpoint=self.load_publish_endpoint, + ) + if resolved_range is not None: + descriptor["load_endpoint_port_base"] = resolved_range[1] + descriptor["load_topic"] = LOAD_TOPIC + return descriptor def should_report_expert_balancedness(self) -> bool: cfg = resolving_view(self) diff --git a/test/registered/unit/disaggregation/test_kv_events.py b/test/registered/unit/disaggregation/test_kv_events.py index 49db8967d..e1bf23478 100644 --- a/test/registered/unit/disaggregation/test_kv_events.py +++ b/test/registered/unit/disaggregation/test_kv_events.py @@ -18,6 +18,7 @@ from sglang.srt.disaggregation.kv_events import ( KVEventBatch, StorageMedium, ZmqEventPublisher, + resolve_load_pub_range, select_kv_publisher_dp_rank, ) from sglang.test.ci.ci_register import register_cpu_ci @@ -26,6 +27,89 @@ from sglang.test.test_utils import CustomTestCase register_cpu_ci(est_time=2, suite="base-a-test-cpu") +class TestResolveLoadPubRange(CustomTestCase): + """The single source of truth both the bind and /server_info route through.""" + + @staticmethod + def _base(kv, replay=None, dp_size=1, explicit="auto"): + resolved, _ = resolve_load_pub_range( + kv_endpoint=kv, + replay_endpoint=replay, + dp_size=dp_size, + load_publish_endpoint=explicit, + ) + return None if resolved is None else resolved[1] + + def test_off_by_default(self): + # Opt-in: unset or "off" reserves nothing, even with a valid config. + self.assertIsNone(self._base("tcp://*:5557", explicit=None)) + self.assertIsNone(self._base("tcp://*:5557", explicit="off")) + + def test_auto_packs_after_kv_range(self): + self.assertEqual(self._base("tcp://*:5557"), 5558) + self.assertEqual(self._base("tcp://*:5557", dp_size=2), 5559) + + def test_auto_skips_an_overlapping_replay_range(self): + # Conventional replay = kv + 1 always overlaps the packed candidate. + self.assertEqual(self._base("tcp://*:5557", "tcp://*:5558"), 5559) + self.assertEqual(self._base("tcp://*:5557", "tcp://*:5558", dp_size=4), 5562) + + def test_non_adjacent_replay_leaves_packing_unchanged(self): + self.assertEqual(self._base("tcp://*:5557", "tcp://*:6000"), 5558) + + def test_auto_declines_connect_style_and_underivable_endpoints(self): + for kv in ( + "tcp://10.0.0.5:5557", # concrete host: connect-style + "tcp://[2001:db8::5]:5557", # concrete IPv6 ("::" is not a wildcard) + "tcp://::1:5557", # bare IPv6: ambiguous + "tcp://host", # no port + "ipc:///tmp/kv", + None, + ): + with self.subTest(kv=kv): + self.assertIsNone(self._base(kv)) + + def test_auto_declines_on_u16_overflow(self): + self.assertIsNone(self._base("tcp://*:65535")) + + def test_explicit_endpoint_moves_and_validates_the_range(self): + self.assertEqual(self._base("tcp://*:5557", explicit="tcp://*:7000"), 7000) + # A concrete explicit host, or one overlapping the kv range, declines. + self.assertIsNone(self._base("tcp://*:5557", explicit="tcp://10.0.0.5:7000")) + self.assertIsNone( + self._base("tcp://*:5557", dp_size=4, explicit="tcp://*:5558") + ) + + def test_reason_is_set_only_for_actionable_declines(self): + # Off by default is unremarkable (no reason); an opt-in the operator + # asked for that can't resolve is worth surfacing. + _, quiet = resolve_load_pub_range( + kv_endpoint="tcp://10.0.0.5:5557", replay_endpoint=None, dp_size=1 + ) + self.assertIsNone(quiet) + _, auto_loud = resolve_load_pub_range( + kv_endpoint="tcp://10.0.0.5:5557", # connect-style: can't derive + replay_endpoint=None, + dp_size=1, + load_publish_endpoint="auto", + ) + self.assertIsNotNone(auto_loud) + # A missing config surfaces at startup — no message may render a bare + # "None". Both the likely mistakes (auto and an explicit address + # without --kv-events-config) go through this. + for endpoint in ("auto", "tcp://*:7000"): + with self.subTest(endpoint=endpoint): + _, no_cfg = resolve_load_pub_range( + kv_endpoint=None, + replay_endpoint=None, + dp_size=1, + load_publish_endpoint=endpoint, + ) + self.assertIsNotNone(no_cfg) + self.assertNotIn("None", no_cfg) + self.assertIn("--kv-events-config", no_cfg) + + class TestSelectKvPublisherDpRank(CustomTestCase): def test_select_rank_across_modes(self): # (label, attn_dp_size, attn_dp_rank, dp_rank, expected) diff --git a/test/registered/unit/entrypoints/test_server_info.py b/test/registered/unit/entrypoints/test_server_info.py index 8aa9b3276..d6a57c61b 100644 --- a/test/registered/unit/entrypoints/test_server_info.py +++ b/test/registered/unit/entrypoints/test_server_info.py @@ -112,6 +112,7 @@ class TestServerInfoKvEventsField(CustomTestCase): kv_events_config=( '{"publisher": "zmq", "endpoint": "tcp://*:5557", "topic": "kv"}' ), + load_publish_endpoint="auto", page_size=64, dp_size=2, ) @@ -128,10 +129,143 @@ class TestServerInfoKvEventsField(CustomTestCase): "topic": "kv", "block_size": 64, "dp_size": 2, + # Load range packed immediately after the KV range: 5557 + dp_size. + "load_endpoint_port_base": 5559, + "load_topic": "load", }, ) + def test_load_port_skips_an_overlapping_replay_range(self): + # Conventional replay = kv + 1: the load range must be advertised + # past the replay ROUTER range (5558 + dp_size), matching where + # SchedulerLoadPublisher actually binds — both sides resolve it via + # resolve_load_pub_range. + args = ServerArgs( + model_path="dummy", + kv_events_config=( + '{"publisher": "zmq", "endpoint": "tcp://*:5557", ' + '"replay_endpoint": "tcp://*:5558"}' + ), + load_publish_endpoint="auto", + page_size=64, + dp_size=2, + ) + + info = _call_server_info_with(args) + + self.assertEqual(info["kv_events"]["load_endpoint_port_base"], 5560) + self.assertEqual(info["kv_events"]["load_topic"], "load") + + def test_explicit_load_publish_endpoint_moves_the_advertised_base(self): + args = ServerArgs( + model_path="dummy", + kv_events_config='{"publisher": "zmq", "endpoint": "tcp://*:5557"}', + load_publish_endpoint="tcp://*:7000", + page_size=64, + dp_size=2, + ) + + info = _call_server_info_with(args) + + self.assertEqual(info["kv_events"]["load_endpoint_port_base"], 7000) + self.assertEqual(info["kv_events"]["load_topic"], "load") + + def test_load_keys_omitted_for_connect_style_kv_endpoint(self): + # A concrete host is connected to rather than bound: the KV-events + # descriptor is still valid (KV events work connect-style), but no + # load range can be bound there, so the load keys must be omitted + # rather than advertising a port nothing listens on. + args = ServerArgs( + model_path="dummy", + kv_events_config=( + '{"publisher": "zmq", "endpoint": "tcp://10.0.0.5:5557"}' + ), + load_publish_endpoint="auto", + page_size=64, + dp_size=1, + ) + + info = _call_server_info_with(args) + + self.assertIsNotNone(info["kv_events"]) + self.assertEqual(info["kv_events"]["endpoint_host"], "10.0.0.5") + self.assertNotIn("load_endpoint_port_base", info["kv_events"]) + self.assertNotIn("load_topic", info["kv_events"]) + + def test_ipv6_wildcard_endpoint_advertises_bracketed_host_and_load_keys(self): + # "[::]" is a bind-all wildcard: the descriptor must keep the + # brackets (consumers splice tcp://{host}:{port}) and the load + # range resolves right after the KV range. + args = ServerArgs( + model_path="dummy", + kv_events_config='{"publisher": "zmq", "endpoint": "tcp://[::]:5557"}', + load_publish_endpoint="auto", + page_size=64, + dp_size=1, + ) + + info = _call_server_info_with(args) + + self.assertEqual(info["kv_events"]["endpoint_host"], "[::]") + self.assertEqual(info["kv_events"]["endpoint_port_base"], 5557) + self.assertEqual(info["kv_events"]["load_endpoint_port_base"], 5558) + + def test_concrete_ipv6_endpoint_advertises_kv_but_not_load(self): + # A concrete IPv6 host works connect-style for KV events (advertised, + # brackets kept) but is not bindable for the load range — and "::" + # appearing inside the address must not be mistaken for a wildcard. + args = ServerArgs( + model_path="dummy", + kv_events_config=( + '{"publisher": "zmq", "endpoint": "tcp://[2001:db8::5]:5557"}' + ), + load_publish_endpoint="auto", + page_size=64, + dp_size=1, + ) + + info = _call_server_info_with(args) + + self.assertEqual(info["kv_events"]["endpoint_host"], "[2001:db8::5]") + self.assertNotIn("load_endpoint_port_base", info["kv_events"]) + + def test_load_keys_omitted_when_explicitly_off(self): + args = ServerArgs( + model_path="dummy", + kv_events_config='{"publisher": "zmq", "endpoint": "tcp://*:5557"}', + load_publish_endpoint="off", + page_size=64, + dp_size=1, + ) + + info = _call_server_info_with(args) + + self.assertIsNotNone(info["kv_events"]) + self.assertNotIn("load_endpoint_port_base", info["kv_events"]) + self.assertNotIn("load_topic", info["kv_events"]) + + def test_load_keys_omitted_when_no_load_range_fits(self): + # kv base 65535 leaves no u16 room for a load range: the kv_events + # descriptor must still be served, with only the load keys omitted, + # so routers fall back to their in-flight counter for load. + args = ServerArgs( + model_path="dummy", + kv_events_config='{"publisher": "zmq", "endpoint": "tcp://*:65535"}', + load_publish_endpoint="auto", + page_size=64, + dp_size=1, + ) + + info = _call_server_info_with(args) + + self.assertIsNotNone(info["kv_events"]) + self.assertEqual(info["kv_events"]["endpoint_port_base"], 65535) + self.assertNotIn("load_endpoint_port_base", info["kv_events"]) + self.assertNotIn("load_topic", info["kv_events"]) + def test_kv_events_descriptor_carries_specific_host_and_topic(self): + # No --load-publish-endpoint: KV descriptor served, load keys absent. + # Upgrade safety rests on this default silence, so pin it here. args = ServerArgs( model_path="dummy", kv_events_config=( @@ -149,6 +283,8 @@ class TestServerInfoKvEventsField(CustomTestCase): self.assertEqual(info["kv_events"]["topic"], "kv") self.assertEqual(info["kv_events"]["block_size"], 128) self.assertEqual(info["kv_events"]["dp_size"], 1) + self.assertNotIn("load_endpoint_port_base", info["kv_events"]) + self.assertNotIn("load_topic", info["kv_events"]) # ----- disabled / unconfigured ------------------------------------- @@ -387,5 +523,66 @@ class TestServerInfoExistingFieldsPreserved(CustomTestCase): json.dumps(info) +class TestLoadPublishEndpointValidation(CustomTestCase): + """--load-publish-endpoint fails fast at the entrypoint, not silently in a + scheduler subprocess log.""" + + def test_requires_kv_events_config(self): + args = ServerArgs(model_path="dummy", load_publish_endpoint="tcp://*:6000") + with self.assertRaisesRegex(ValueError, "kv-events"): + args.check_load_publish_args() + + def test_rejects_non_bindable_endpoint(self): + args = ServerArgs( + model_path="dummy", + kv_events_config='{"publisher": "zmq", "endpoint": "tcp://*:5557"}', + load_publish_endpoint="tcp://10.0.0.5:6000", + ) + with self.assertRaisesRegex(ValueError, "bindable"): + args.check_load_publish_args() + + def test_rejects_endpoint_overlapping_the_kv_range(self): + args = ServerArgs( + model_path="dummy", + kv_events_config='{"publisher": "zmq", "endpoint": "tcp://*:5557"}', + dp_size=4, + load_publish_endpoint="tcp://*:5558", + ) + with self.assertRaisesRegex(ValueError, "overlaps"): + args.check_load_publish_args() + + def test_rejects_null_publisher(self): + # publisher='null' disables KV events, so there is nothing to advertise + # through — accepting the opt-in would silently do nothing. + args = ServerArgs( + model_path="dummy", + kv_events_config='{"publisher": "null"}', + load_publish_endpoint="auto", + ) + with self.assertRaisesRegex(ValueError, "null"): + args.check_load_publish_args() + + def test_rejects_unparseable_kv_events_config(self): + args = ServerArgs( + model_path="dummy", + kv_events_config="{not json", + load_publish_endpoint="auto", + ) + with self.assertRaisesRegex(ValueError, "not parseable"): + args.check_load_publish_args() + + def test_off_and_valid_endpoint_pass(self): + for endpoint in (None, "off", "OFF", "auto", "tcp://*:6000"): + with self.subTest(endpoint=endpoint): + args = ServerArgs( + model_path="dummy", + kv_events_config=( + '{"publisher": "zmq", "endpoint": "tcp://*:5557"}' + ), + load_publish_endpoint=endpoint, + ) + args.check_load_publish_args() # must not raise + + if __name__ == "__main__": unittest.main() diff --git a/test/registered/unit/managers/test_loadstat_wire.py b/test/registered/unit/managers/test_loadstat_wire.py new file mode 100644 index 000000000..3c392aeb9 --- /dev/null +++ b/test/registered/unit/managers/test_loadstat_wire.py @@ -0,0 +1,502 @@ +"""Wire contract and port/rank gating for the LoadStat load snapshot. + +Locks the msgpack array shape the sgl-router `cache_aware_zmq` policy will +decode positionally (that consumer lands with the router PR; it is not yet +in this tree, so this pins only the Python side): + + ["LoadStat", num_running_reqs, num_waiting_reqs, num_tokens, + max_total_num_tokens, attn_dp_rank] + +carried as the payload of a three-frame message ``[b"load", BE-i64 seq, +payload]``. A field reorder or rename is a silent cross-language break, so +`test_loadstat_golden_bytes` pins the exact encoding — assert the same hex +on the Rust side when that PR lands to actually close the loop. +TestLoadPublisherGating pins which schedulers publish and on which port. +CPU-only: the socket bind is stubbed at the `_open_pub_socket` seam. +""" + +import unittest +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +import msgspec.msgpack + +from sglang.srt.distributed.parallel_state_wrapper import ParallelState +from sglang.srt.managers.scheduler_components.load_publisher import ( + LoadStat, + SchedulerLoadPublisher, +) +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=2, suite="base-a-test-cpu") + + +class TestLoadStatWire(CustomTestCase): + def test_loadstat_golden_bytes(self): + # Exact on-the-wire encoding. Assert the identical hex from the Rust + # decoder's test when the router PR lands — that is what actually pins + # a cross-language format; the decode-round-trip below only pins Python. + raw = msgspec.msgpack.Encoder().encode( + LoadStat( + num_running_reqs=7, + num_waiting_reqs=3, + num_tokens=1024, + max_total_num_tokens=8192, + attn_dp_rank=2, + ) + ) + self.assertEqual(raw.hex(), "96a84c6f6164537461740703cd0400cd200002") + + def test_loadstat_msgpack_array_shape(self): + raw = msgspec.msgpack.Encoder().encode( + LoadStat( + num_running_reqs=7, + num_waiting_reqs=3, + num_tokens=1024, + max_total_num_tokens=8192, + attn_dp_rank=2, + ) + ) + # tag=True + array_like → [tag, *fields] in declaration order; the + # router reads the tag + four counts and ignores the trailing field. + self.assertEqual( + msgspec.msgpack.Decoder().decode(raw), + ["LoadStat", 7, 3, 1024, 8192, 2], + ) + + def test_loadstat_tag_is_class_name(self): + # The tag is the literal "LoadStat"; guard against an accidental + # msgspec `tag=` override or a class rename. + raw = msgspec.msgpack.Encoder().encode( + LoadStat( + num_running_reqs=0, + num_waiting_reqs=0, + num_tokens=0, + max_total_num_tokens=0, + ) + ) + decoded = msgspec.msgpack.Decoder().decode(raw) + # LoadStat sets no omit_defaults, so the trailing field is always + # emitted (null when unset); a decoder must tolerate it. + self.assertEqual(decoded, ["LoadStat", 0, 0, 0, 0, None]) + + +ZMQ_ENDPOINT = '{"publisher": "zmq", "endpoint": "tcp://*:5557"}' + + +class TestLoadPublisherGating(CustomTestCase): + """One load publisher per independent KV cache, on a resolvable port. + + Getting either half wrong makes several schedulers bind the same port, + which is an uncaught ZMQError at startup for a bind-style endpoint and — + worse — silently merges every worker's load onto one rank for a + connect-style one. + """ + + def _build( + self, *, config=ZMQ_ENDPOINT, dp_size=1, explicit="auto", **ps_overrides + ): + """Construct a publisher with the socket bind stubbed out, returning + (publisher, captured _open_pub_socket mock). Opts in via explicit="auto" + by default (the feature is off without it). dp_size lives on the ps, + which the publisher reads (no separate param to disagree with it).""" + with patch( + "sglang.srt.managers.scheduler_components.load_publisher." + "_open_pub_socket" + ) as open_sock: + pub = SchedulerLoadPublisher( + kv_events_config=config, + ps=ParallelState.trivial(dp_size=dp_size, **ps_overrides), + load_publish_endpoint=explicit, + ) + return pub, open_sock + + def test_disabled_by_default(self): + # Off unless opted in: a bare --kv-events-config user (no + # --load-publish-endpoint) reserves no load port, so an upgrade can't + # collide with a co-hosted neighbor's KV bind. + pub, open_sock = self._build(explicit=None) + self.assertFalse(pub.enable) + open_sock.assert_not_called() + + def test_enabled_on_rank_zero(self): + pub, open_sock = self._build() # explicit="auto" + self.assertTrue(pub.enable) + open_sock.assert_called_once_with("tcp://*:5558") + + def test_disabled_off_pp_rank_zero(self): + # Every PP stage shares attn_tp_rank/attn_cp_rank 0, so without the + # pp_rank gate they all bind the same load port. + pub, open_sock = self._build(pp_rank=1, pp_size=2) + self.assertFalse(pub.enable) + open_sock.assert_not_called() + + def test_disabled_off_attn_tp_and_cp_rank_zero(self): + for override in ({"attn_tp_rank": 1}, {"attn_cp_rank": 1}): + with self.subTest(**override): + pub, open_sock = self._build(**override) + self.assertFalse(pub.enable) + open_sock.assert_not_called() + + def test_pure_dp_keys_the_load_port_by_dp_rank(self): + # Pure DP: attn_dp_size == 1 and every worker has attn_dp_rank == 0, so + # the publisher must key off dp_rank or all replicas collide on one + # port. kv 5557 + dp_size 4 => base 5561; rank 2 binds 5563. + _, open_sock = self._build(attn_dp_size=1, attn_dp_rank=0, dp_rank=2, dp_size=4) + open_sock.assert_called_once_with("tcp://*:5563") + + def test_dp_attention_keys_the_load_port_by_attn_dp_rank(self): + _, open_sock = self._build(attn_dp_size=4, attn_dp_rank=3, dp_rank=0, dp_size=4) + open_sock.assert_called_once_with("tcp://*:5564") + + def test_load_port_is_packed_after_the_kv_range(self): + _, open_sock = self._build(dp_size=2) + open_sock.assert_called_once_with("tcp://*:5559") + + def test_accepts_every_bind_style_host(self): + for host, expected in ( + ("*", "tcp://*:5558"), + ("0.0.0.0", "tcp://0.0.0.0:5558"), + ("[::]", "tcp://[::]:5558"), + ): + with self.subTest(host=host): + pub, open_sock = self._build( + config='{"publisher": "zmq", "endpoint": "tcp://%s:5557"}' % host + ) + self.assertTrue(pub.enable) + open_sock.assert_called_once_with(expected) + + def test_unresolvable_endpoint_declines_instead_of_raising(self): + # ipc:// and inproc:// are valid KV-event endpoints but carry no port + # to pack after; port-less/malformed tcp shapes are underivable; and a + # concrete host (IPv4 or IPv6 — "::" appears inside every IPv6 + # address, so this must not be a substring test) would be *connected + # to* rather than bound, publishing into a void. None of them may + # take down scheduler startup over a load socket. + for endpoint in ( + "ipc:///tmp/kv.sock", + "inproc://kv", + "tcp://somehost", + "tcp://*:*", + "tcp://somehost:-100", + "tcp://10.0.0.5:5557", + "tcp://[2001:db8::5]:5557", + "tcp://::1:5557", + ): + with self.subTest(endpoint=endpoint): + pub, open_sock = self._build( + config='{"publisher": "zmq", "endpoint": "%s"}' % endpoint + ) + self.assertFalse(pub.enable) + open_sock.assert_not_called() + + def test_disabled_paths_leave_the_publisher_unbound(self): + # Every bail-out must leave the socket unbound (surfaced as + # enable == False) so publish_load_stat returns before computing + # the (non-trivial) load snapshot. + for label, config in ( + ("no config", None), + ("null publisher", '{"publisher": "null"}'), + ("malformed", "{not json"), + ): + with self.subTest(label): + pub, _ = self._build(config=config) + self.assertFalse(pub.enable) + + def test_replay_port_collision_skips_past_the_replay_range(self): + # Conventional config inherited from upstream: KV on 5557, replay on + # 5558. With dp_size=1 the load socket would land exactly on the + # replay ROUTER's port; instead of declining (which would silently + # turn the feature off on exactly this common config) the load range + # packs after the replay range: 5558 + dp_size = 5559. + pub, open_sock = self._build( + config='{"publisher": "zmq", "endpoint": "tcp://*:5557", ' + '"replay_endpoint": "tcp://*:5558"}' + ) + self.assertTrue(pub.enable) + open_sock.assert_called_once_with("tcp://*:5559") + + def test_replay_skip_covers_the_whole_per_rank_range(self): + # dp_size=4: KV range 5557..5560, replay ROUTER range 5558..5561; the + # first candidate (5561) still collides with the replay range's tail, + # so the load range packs after it: 5558 + 4 = 5562. + pub, open_sock = self._build( + config='{"publisher": "zmq", "endpoint": "tcp://*:5557", ' + '"replay_endpoint": "tcp://*:5558"}', + dp_size=4, + ) + self.assertTrue(pub.enable) + open_sock.assert_called_once_with("tcp://*:5562") + + def test_replay_far_away_keeps_the_packed_port(self): + # No overlap with the replay range => the load range stays right + # after the KV range (no needless jump past a distant replay port). + pub, open_sock = self._build( + config='{"publisher": "zmq", "endpoint": "tcp://*:5557", ' + '"replay_endpoint": "tcp://*:6000"}' + ) + self.assertTrue(pub.enable) + open_sock.assert_called_once_with("tcp://*:5558") + + def test_port_overflow_declines_instead_of_crashing(self): + # kv base 65535 + dp_size pushes the load range past u16; + # /server_info omits the key for the same reason. + pub, open_sock = self._build( + config='{"publisher": "zmq", "endpoint": "tcp://*:65535"}' + ) + self.assertFalse(pub.enable) + open_sock.assert_not_called() + + def test_explicit_endpoint_moves_the_range(self): + # --load-publish-endpoint sets the range outright; rank r still binds + # base + r (pure DP keys by dp_rank). + pub, open_sock = self._build(explicit="tcp://*:7000") + self.assertTrue(pub.enable) + open_sock.assert_called_once_with("tcp://*:7000") + + _, open_sock = self._build( + explicit="tcp://*:7000", + attn_dp_size=1, + attn_dp_rank=0, + dp_rank=2, + dp_size=4, + ) + open_sock.assert_called_once_with("tcp://*:7002") + + def test_explicit_endpoint_must_be_bindable(self): + # A concrete host would be connected to rather than bound. + pub, open_sock = self._build(explicit="tcp://10.0.0.5:7000") + self.assertFalse(pub.enable) + open_sock.assert_not_called() + + def test_explicit_off_disables_load_publishing(self): + # The operator's off switch: KV events without the extra port range. + # /server_info omits the load keys through the same resolver. + pub, open_sock = self._build(explicit="off") + self.assertFalse(pub.enable) + open_sock.assert_not_called() + + def test_bind_failure_disables_without_raising(self): + # An occupied port must not take down scheduler startup over a routing + # hint; the publisher logs and stays a no-op. Opted in (auto) so the + # bind is actually reached — otherwise the feature is just off. + import zmq + + with patch( + "sglang.srt.managers.scheduler_components.load_publisher." + "_open_pub_socket", + side_effect=zmq.ZMQError, + ) as open_sock: + pub = SchedulerLoadPublisher( + kv_events_config=ZMQ_ENDPOINT, + ps=ParallelState.trivial(), + load_publish_endpoint="auto", + ) + open_sock.assert_called_once() # the bind was attempted and failed + self.assertFalse(pub.enable) + pub.publish_load_stat(MagicMock(), force=True) # still a no-op + + def test_close_is_idempotent_and_disables(self): + pub, _ = self._build() + socket = pub._socket + pub.close() + pub.close() + socket.close.assert_called_once() + self.assertFalse(pub.enable) + provider = MagicMock() + pub.publish_load_stat(provider, force=True) + provider.assert_not_called() + + def test_explicit_endpoint_needs_an_advertisable_kv_endpoint(self): + # Discovery rides on /server_info's kv_events block, which is absent + # for non-tcp (or port-less) KV endpoints — binding the explicit + # range anyway would claim a port no router can ever find. + for kv_endpoint in ("ipc:///tmp/kv.sock", "inproc://kv", "tcp://0.0.0.0"): + with self.subTest(kv_endpoint=kv_endpoint): + pub, open_sock = self._build( + config='{"publisher": "zmq", "endpoint": "%s"}' % kv_endpoint, + explicit="tcp://*:7000", + ) + self.assertFalse(pub.enable) + open_sock.assert_not_called() + + def test_explicit_endpoint_inside_the_kv_range_declines(self): + # The KV publisher binds its own range later and unguarded, so taking + # one of its ports would kill startup blaming the KV publisher. + pub, open_sock = self._build(dp_size=4, explicit="tcp://*:5558") + self.assertFalse(pub.enable) + open_sock.assert_not_called() + + # ----- publish path ------------------------------------------------- + + @staticmethod + def _provider(running): + return MagicMock( + return_value=SimpleNamespace( + num_running_reqs=running, + num_waiting_reqs=2, + num_used_tokens=3, + max_total_num_tokens=4, + ) + ) + + def test_publish_skips_snapshot_when_disabled(self): + pub, _ = self._build(config='{"publisher": "null"}') + provider = MagicMock() + pub.publish_load_stat(provider, force=True) + provider.assert_not_called() + + def test_caller_supplied_snapshot_bypasses_the_provider(self): + # The scheduler hands in the snapshot it already computed for the + # DP-balancing sink; the provider is the fallback for cycles where + # that sink was throttled — it must not run when a snapshot is given. + pub, _ = self._build() + provider = MagicMock() + snap = SimpleNamespace( + num_running_reqs=1, + num_waiting_reqs=2, + num_used_tokens=3, + max_total_num_tokens=4, + ) + pub.publish_load_stat(provider, force=True, snapshot=snap) + provider.assert_not_called() + self.assertEqual(pub._socket.send_multipart.call_count, 1) + + def test_publish_frames_are_topic_seq_payload(self): + # Three frames, matching the KV-event socket's layout so one + # subscriber loop handles both. + pub, _ = self._build() + pub.publish_load_stat(self._provider(running=1), force=True) + (frames,), _ = pub._socket.send_multipart.call_args + topic, seq, payload = frames + self.assertEqual(topic, b"load") + self.assertEqual(seq, (0).to_bytes(8, "big")) + self.assertEqual( + msgspec.msgpack.Decoder().decode(payload), + ["LoadStat", 1, 2, 3, 4, 0], + ) + + def test_unchanged_stat_is_deduped_to_the_heartbeat(self): + # force=True fires per idle-loop iteration (which busy-spins without + # --sleep-on-idle); an unchanged gauge must go out once per heartbeat, + # not per iteration. time is patched so the test cannot race the + # wall clock. + pub, _ = self._build() + provider = self._provider(running=1) + with patch( + "sglang.srt.managers.scheduler_components.load_publisher.time" + ) as fake_time: + fake_time.monotonic.return_value = 100.0 + pub.publish_load_stat(provider, force=True) # first: publishes + pub.publish_load_stat(provider, force=True) # unchanged: deduped + self.assertEqual(pub._socket.send_multipart.call_count, 1) + fake_time.monotonic.return_value = 101.5 # heartbeat elapsed + pub.publish_load_stat(provider, force=True) + self.assertEqual(pub._socket.send_multipart.call_count, 2) + + def test_call_throttle_stays_engaged_across_dedup_hits(self): + # Regression: the counter must reset when the throttle PASSES, not + # when a send happens. Resetting only on the send path let one dedup + # hit saturate the counter, running the O(queue) provider every step. + # A working counter fires the provider at counts 5 and 10. + pub, _ = self._build() + provider = self._provider(running=1) + with patch( + "sglang.srt.managers.scheduler_components.load_publisher.time" + ) as fake_time: + fake_time.monotonic.return_value = 100.0 + for _ in range(10): + pub.publish_load_stat(provider) + self.assertEqual(provider.call_count, 2) + + def test_provider_failure_never_raises(self): + # get_loads raising must not crash the scheduler loop; it warns and + # leaves the counter reset (not saturated). + def boom(): + raise RuntimeError("get_loads exploded") + + pub, _ = self._build() + with self.assertLogs( + "sglang.srt.managers.scheduler_components.load_publisher", + level="WARNING", + ): + pub.publish_load_stat(boom, force=True) + self.assertEqual(pub._publish_counter, 0) + + def test_changed_stat_publishes_immediately(self): + # The busy->idle (and idle->busy) transition must never be delayed: + # a changed gauge bypasses the heartbeat dedup even when the last + # send was a moment ago. + pub, _ = self._build() + with patch( + "sglang.srt.managers.scheduler_components.load_publisher.time" + ) as fake_time: + fake_time.monotonic.return_value = 100.0 + pub.publish_load_stat(self._provider(running=7), force=True) + pub.publish_load_stat(self._provider(running=0), force=True) + self.assertEqual(pub._socket.send_multipart.call_count, 2) + + +class TestLoadStatIntegration(CustomTestCase): + """The one path every gating test stubs: a real socket bind + SUB + round-trip. Covers _open_pub_socket (bind, HWM/LINGER/IPV6 order) and the + three-frame wire end to end.""" + + def test_binds_and_delivers_three_decodable_frames(self): + import socket as _socket + import time as _time + + import zmq + + # Probe on "" (all interfaces) to match ZMQ's wildcard bind, and retry: + # probe-then-bind is a TOCTOU race and the publisher swallows bind + # errors, so a lost race shows up only as a disabled publisher. + pub = None + for _ in range(3): + with _socket.socket() as probe: + probe.bind(("", 0)) + port = probe.getsockname()[1] + pub = SchedulerLoadPublisher( + kv_events_config='{"publisher": "zmq", "endpoint": "tcp://*:5557"}', + ps=ParallelState.trivial(), + load_publish_endpoint=f"tcp://*:{port}", + ) + if pub.enable: + break + self.assertTrue(pub.enable, "load socket never bound a free port") + self.addCleanup(pub.close) + + sub = zmq.Context.instance().socket(zmq.SUB) + sub.connect(f"tcp://127.0.0.1:{port}") + sub.setsockopt_string(zmq.SUBSCRIBE, "load") # exact advertised topic + self.addCleanup(sub.close) + + snap = SimpleNamespace( + num_running_reqs=7, + num_waiting_reqs=3, + num_used_tokens=1024, + max_total_num_tokens=8192, + ) + # PUB/SUB drops messages sent before the subscription propagates, so + # re-publish until one lands (heartbeat reset each pass). + frames = None + deadline = _time.time() + 5 + while frames is None and _time.time() < deadline: + pub._last_publish_ts = 0.0 + pub.publish_load_stat(lambda: snap, force=True) + if sub.poll(100): + frames = sub.recv_multipart() + self.assertIsNotNone(frames, "no load frame received within 5s") + + topic, seq, payload = frames + self.assertEqual(topic, b"load") + self.assertEqual(len(seq), 8) + self.assertEqual( + msgspec.msgpack.Decoder().decode(payload), + ["LoadStat", 7, 3, 1024, 8192, 0], + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/managers/test_scheduler_on_idle_load.py b/test/registered/unit/managers/test_scheduler_on_idle_load.py new file mode 100644 index 000000000..205ca24da --- /dev/null +++ b/test/registered/unit/managers/test_scheduler_on_idle_load.py @@ -0,0 +1,52 @@ +"""on_idle's stalled-path load publish is wall-clock bounded. + +A no-batch-but-not-idle stall spins on_idle without sleeping, so the gate must +cap the O(queue) get_loads for both the DP-balancing writer and the load +socket. CPU-only: builds a bare Scheduler with mocked collaborators, like +test_scheduler_flush_cache. +""" + +import unittest +from unittest.mock import MagicMock, patch + +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase, maybe_stub_sgl_kernel + +maybe_stub_sgl_kernel() + +from sglang.srt.managers.scheduler import Scheduler + +register_cpu_ci(est_time=2, suite="base-a-test-cpu") + + +class TestOnIdleStallPublish(CustomTestCase): + def _stalled_scheduler(self) -> Scheduler: + s = Scheduler.__new__(Scheduler) + s.maybe_send_health_check_signal = MagicMock() + s.is_fully_idle = MagicMock(return_value=False) # stalled, not idle + s.publish_load_snapshot = MagicMock(return_value=None) + s.load_publisher = MagicMock() + s.load_inquirer = MagicMock() + s._last_stall_publish_ts = float("-inf") + return s + + def test_spinning_stall_publishes_once_within_the_floor(self): + s = self._stalled_scheduler() + with patch("sglang.srt.managers.scheduler.time.monotonic", return_value=100.0): + for _ in range(100): + s.on_idle() + self.assertEqual(s.publish_load_snapshot.call_count, 1) + self.assertEqual(s.load_publisher.publish_load_stat.call_count, 1) + + def test_publishes_again_after_the_floor_elapses(self): + s = self._stalled_scheduler() + with patch("sglang.srt.managers.scheduler.time.monotonic") as mono: + mono.return_value = 100.0 + s.on_idle() + mono.return_value = 100.10 # > LOAD_STALL_REFRESH_S + s.on_idle() + self.assertEqual(s.publish_load_snapshot.call_count, 2) + + +if __name__ == "__main__": + unittest.main()