diff --git a/python/sglang/srt/arg_groups/fields/disagg.py b/python/sglang/srt/arg_groups/fields/disagg.py index 8414e4894..ad80c251a 100644 --- a/python/sglang/srt/arg_groups/fields/disagg.py +++ b/python/sglang/srt/arg_groups/fields/disagg.py @@ -114,6 +114,10 @@ class Disagg(msgspec.Struct): int, "The interval to poll requests in decode server. Can be set to >1 to reduce the overhead of this.", ] = 1 + enable_pd_role_switch: A[ + bool, + "Allow runtime prefill<->decode role switch via /pd_role_switch (PD mode).", + ] = False optimistic_prefill_attempts: A[ int, "Number of optimistic prefill forward passes that skip the bootstrap wait." ] = 0 diff --git a/python/sglang/srt/arg_groups/pd_disaggregation_hook.py b/python/sglang/srt/arg_groups/pd_disaggregation_hook.py index f66449f3f..6fa2ea78f 100644 --- a/python/sglang/srt/arg_groups/pd_disaggregation_hook.py +++ b/python/sglang/srt/arg_groups/pd_disaggregation_hook.py @@ -147,6 +147,39 @@ def handle_pd_disaggregation(server_args: ServerArgs) -> None: f"got '{cfg.disaggregation_transfer_backend}'." ) + # Reject features whose role-specific state is not rebuilt on a flip. + if cfg.enable_pd_role_switch: + view = resolved_view(server_args) + unsupported = [] + if view.enable_dp_attention: + unsupported.append("DP attention (--enable-dp-attention)") + if view.ep_size > 1: + unsupported.append(f"expert parallelism (--ep-size {view.ep_size})") + if view.moe_a2a_backend != "none": + unsupported.append( + f"MoE all-to-all (--moe-a2a-backend {view.moe_a2a_backend})" + ) + if view.pp_size > 1: + unsupported.append(f"pipeline parallelism (--pp-size {view.pp_size})") + if view.dp_size > 1: + unsupported.append(f"data parallelism (--dp-size {view.dp_size})") + if view.dcp_size > 1: + unsupported.append( + f"decode context parallelism (--dcp-size {view.dcp_size})" + ) + if view.speculative_algorithm is not None: + unsupported.append( + "speculative decoding " + f"(--speculative-algorithm {view.speculative_algorithm})" + ) + if unsupported: + raise ValueError( + "--enable-pd-role-switch does not rebuild role-specific " + "state for the following features: " + + ", ".join(unsupported) + + ". Remove these options or drop --enable-pd-role-switch." + ) + def _alias_bootstrap_port_to_api_port(server_args: ServerArgs) -> None: """Rust-server prefill serves the KV bootstrap registry on the api listener diff --git a/python/sglang/srt/disaggregation/base/conn.py b/python/sglang/srt/disaggregation/base/conn.py index 239504e88..fc93da096 100644 --- a/python/sglang/srt/disaggregation/base/conn.py +++ b/python/sglang/srt/disaggregation/base/conn.py @@ -127,6 +127,15 @@ class BaseKVManager(ABC): """Register prefill server info to the bootstrap server.""" ... + # Opt-in per backend: set True and implement teardown() to support runtime PD + # role switch (release transfer resources; the scheduler owns the KV pool). + supports_role_switch: bool = False + + def teardown(self) -> None: + raise NotImplementedError( + f"{type(self).__name__} does not support PD role switch teardown" + ) + class BaseKVSender(ABC): @abstractmethod diff --git a/python/sglang/srt/disaggregation/common/conn.py b/python/sglang/srt/disaggregation/common/conn.py index 78fe9d585..d37dce6b4 100644 --- a/python/sglang/srt/disaggregation/common/conn.py +++ b/python/sglang/srt/disaggregation/common/conn.py @@ -287,6 +287,9 @@ class CommonKVManager(BaseKVManager): self.max_failures = max( envs.SGLANG_DISAGGREGATION_HEARTBEAT_MAX_FAILURE.get(), 1 ) + # Event used to signal the heartbeat checker thread to exit + # during teardown (e.g. runtime P<->D role switch). + self._heartbeat_shutdown = threading.Event() # If a timeout happens on the decode side, it means decode instances # fail to receive the KV Cache transfer done signal after bootstrapping. # These timeout requests should be aborted to release the tree cache. @@ -707,6 +710,25 @@ class CommonKVManager(BaseKVManager): return self._kv_replica_factor = info.required_dst_info_num + def _make_worker_recv(self, socket, timeout_ms: int = 500): + """Build the blocking multipart recv used by a worker thread. + + Plain blocking recv unless role switching is enabled: teardown flips a + stop flag that a blocked recv can never observe, so in that mode poll + with a timeout and return None when it expires. Deployments without + --enable-pd-role-switch keep the original blocking recv and pay nothing. + """ + if not self.server_args.enable_pd_role_switch: + return socket.recv_multipart + + poller = zmq.Poller() + poller.register(socket, zmq.POLLIN) + + def recv(): + return socket.recv_multipart() if poller.poll(timeout_ms) else None + + return recv + def _ensure_prefill_recompute_executor( self, ) -> concurrent.futures.ThreadPoolExecutor: @@ -1308,12 +1330,18 @@ class CommonKVManager(BaseKVManager): return src_kv_ptrs, sliced_dst - def _start_heartbeat_checker_thread(self): - """Start the heartbeat checker thread for Decode worker.""" + def _start_heartbeat_checker_thread(self) -> threading.Thread: + """Start the heartbeat checker thread for Decode worker. + + Returns the thread object so callers can track/join it during teardown. + """ def heartbeat_checker(): - while True: - time.sleep(self.heartbeat_interval) + while not self._heartbeat_shutdown.is_set(): + # Use Event.wait() instead of time.sleep() so teardown can + # wake this thread immediately by setting the event. + if self._heartbeat_shutdown.wait(self.heartbeat_interval): + break with self.connection_lock: addresses = list(self.prefill_info_table.keys()) @@ -1352,7 +1380,13 @@ class CommonKVManager(BaseKVManager): if bootstrap_addr in self.session_pool: del self.session_pool[bootstrap_addr] - threading.Thread(target=heartbeat_checker, daemon=True).start() + t = threading.Thread( + target=heartbeat_checker, + name="HeartbeatChecker", + daemon=True, + ) + t.start() + return t def _on_heartbeat_success(self, bootstrap_addr: str): """Hook called on successful heartbeat. Override for backend-specific cleanup.""" @@ -1800,6 +1834,31 @@ class CommonKVReceiver(BaseKVReceiver): sock.close() logger.debug(f"Disconnected stale ZMQ PUSH socket (receiver): {endpoint}") + @classmethod + def close_all_sockets(cls): + """Close all cached PUSH sockets on role switch, keeping ``_ctx`` reusable.""" + with cls._global_lock: + entries = list(cls._socket_cache.items()) + locks = cls._socket_locks.copy() + cls._socket_cache.clear() + cls._socket_locks.clear() + + # Close outside _global_lock: _connect drops it before the per-endpoint lock. + for endpoint, sock in entries: + lock = locks.get(endpoint) + try: + if lock: + with lock: + sock.close(linger=0) + else: + sock.close(linger=0) + except Exception: + logger.exception( + f"Failed to close ZMQ PUSH socket (receiver): {endpoint}" + ) + if entries: + logger.debug(f"Closed {len(entries)} receiver ZMQ PUSH socket(s)") + @classmethod def _connect_to_bootstrap_server(cls, bootstrap_info: dict): ip_address = bootstrap_info["rank_ip"] diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py index d6bc4aed9..2c3b2b242 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -208,6 +208,8 @@ class KVArgsRegisterInfo: class MooncakeKVManager(StagingManagerMixin, CommonKVManager): AUX_DATA_HEADER = b"AUX_DATA" + # Implements teardown() below, so runtime PD role switching is supported. + supports_role_switch = True def __init__( self, @@ -224,6 +226,9 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager): envs.SGLANG_MOONCAKE_MAX_TRANSFER_BATCH_INDICES.get() ) self.enable_trace = get_observability().enable_trace + # Set by teardown() to make worker threads exit (P<->D role switch). + self._stopped = False + self._worker_threads: List[threading.Thread] = [] if self.disaggregation_mode == DisaggregationMode.PREFILL: self.session_failures = defaultdict(int) self.failed_sessions = set() @@ -262,7 +267,10 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager): for i, (queue, executor) in enumerate( zip(self.transfer_queues, self.executors) ): - threading.Thread( + # Track the thread so teardown() can join it: otherwise every + # P->D->P flip that re-enters PREFILL leaks threads + # (each parked forever in FastQueue.get()). + t = threading.Thread( target=self.transfer_worker, args=( queue, @@ -275,7 +283,9 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager): i, ), daemon=True, - ).start() + ) + t.start() + self._worker_threads.append(t) self.enable_failed_session_probe = ( envs.SGLANG_ENABLE_FAILED_SESSION_PROBE.get() ) @@ -284,11 +294,13 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager): envs.SGLANG_FAILED_SESSION_PROBE_INTERVAL_S.get() ) self._failed_session_probe_shutdown = threading.Event() - threading.Thread( + t = threading.Thread( target=self._failed_session_probe_loop, name="MooncakeFailedSessionProbe", daemon=True, - ).start() + ) + t.start() + self._worker_threads.append(t) elif self.disaggregation_mode == DisaggregationMode.DECODE: self._staging_ctx = DecodeStagingContext() if self.enable_staging else None if self.enable_staging: @@ -338,6 +350,94 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager): with self.connection_lock: self.connection_pool.clear() + def teardown(self) -> None: + """Stop worker threads and release transport resources so this + KVManager can be discarded during a P<->D role switch. + + The KV cache pool memory is owned by the scheduler and is NOT freed + here; only mooncake-side registrations / sockets are released. + """ + self._stopped = True + + # Stop the failed-session probe loop (PREFILL role) if running. + probe_shutdown = getattr(self, "_failed_session_probe_shutdown", None) + if probe_shutdown is not None: + probe_shutdown.set() + + # Stop the heartbeat checker thread (DECODE role) if running. + heartbeat_shutdown = getattr(self, "_heartbeat_shutdown", None) + if heartbeat_shutdown is not None: + heartbeat_shutdown.set() + + # Transfer workers (PREFILL role) park in FastQueue.get(), which has no + # timeout; push a None sentinel per shard to wake and stop them so the + # join below returns instead of leaking the thread. + for queue in getattr(self, "transfer_queues", []): + try: + queue.put(None) + except Exception: + logger.exception( + "Failed to signal mooncake transfer worker on teardown" + ) + + # Shutdown thread pool executors (PREFILL role). + for executor in getattr(self, "executors", []): + try: + executor.shutdown(wait=False, cancel_futures=True) + except TypeError: + # Python < 3.9 does not support cancel_futures + executor.shutdown(wait=False) + except Exception: + logger.exception("Failed to shutdown executor on teardown") + self.executors = [] + + # Join workers before touching their sockets: ZMQ sockets aren't + # thread-safe, so don't close server_socket while a worker may poll it. + for t in self._worker_threads: + t.join(timeout=3.0) + self._worker_threads = [] + + # Drop the queues so their buffered tasks/senders are released too. + self.transfer_queues = [] + + # Close cached PUSH sockets (used by _connect for status sync). + with self._socket_lock: + for sock in self._socket_cache.values(): + try: + sock.close(linger=0) + except Exception: + pass + for monitor in self._monitor_cache.values(): + try: + monitor.close() + except Exception: + pass + self._socket_cache.clear() + self._monitor_cache.clear() + + try: + self.server_socket.close(linger=0) + except Exception: + logger.exception("Failed to close mooncake server_socket during teardown") + + # destroy() force-closes every socket in the context; plain term() + # would block waiting on them. + try: + self._zmq_ctx.destroy(linger=0) + except Exception: + logger.exception("Failed to destroy mooncake zmq context during teardown") + + # Deregister memory from the transfer engine. + try: + self.deregister_buffer_to_engine() + except Exception: + logger.exception("Failed to deregister buffers during teardown") + + logger.info( + "MooncakeKVManager torn down (was role=%s)", + self.disaggregation_mode.value, + ) + # ------------------------------------------------------------------ # Staging buffer methods (all delegate to staging_handler.py) # ------------------------------------------------------------------ @@ -1862,6 +1962,11 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager): while True: try: kv_chunk: TransferKVChunk = queue.get() + # teardown() pushes a None sentinel to unblock get() and stop + # the worker: FastQueue.get() blocks indefinitely, so checking + # _stopped alone can never wake a parked worker during a role switch. + if kv_chunk is None: + break if self.enable_trace: kv_chunk.trace_ctx.rebuild_thread_context() kv_chunk.trace_ctx.trace_slice_start( @@ -2186,11 +2291,15 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager): ) def start_prefill_thread(self): + recv = self._make_worker_recv(self.server_socket) + def bootstrap_thread(): """This thread recvs pre-alloc notification from the decode engine""" # KVPoll.Bootstrapping -> KVPoll.WaitingForInput - while True: - waiting_req_bytes = self.server_socket.recv_multipart() + while not self._stopped: + waiting_req_bytes = recv() + if waiting_req_bytes is None: + continue room = waiting_req_bytes[0].decode("ascii") # Staging: decode reports consumption watermark back to prefill if room == "WATERMARK": @@ -2321,12 +2430,18 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager): ) self.update_status(room, KVPoll.WaitingForInput) - threading.Thread(target=bootstrap_thread).start() + t = threading.Thread(target=bootstrap_thread, daemon=True) + t.start() + self._worker_threads.append(t) def start_decode_thread(self): + recv = self._make_worker_recv(self.server_socket) + def decode_thread(): - while True: - msg = self.server_socket.recv_multipart() + while not self._stopped: + msg = recv() + if msg is None: + continue if msg[0] == MooncakeKVManager.AUX_DATA_HEADER: self._handle_aux_data(msg) continue @@ -2379,8 +2494,11 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager): failure_reason=reason, ) - threading.Thread(target=decode_thread).start() - self._start_heartbeat_checker_thread() + t = threading.Thread(target=decode_thread, daemon=True) + t.start() + self._worker_threads.append(t) + t = self._start_heartbeat_checker_thread() + self._worker_threads.append(t) def add_transfer_request( self, diff --git a/python/sglang/srt/disaggregation/mori/conn.py b/python/sglang/srt/disaggregation/mori/conn.py index 2b62d97ab..6886dac29 100644 --- a/python/sglang/srt/disaggregation/mori/conn.py +++ b/python/sglang/srt/disaggregation/mori/conn.py @@ -293,6 +293,8 @@ class TransferTarget: class MoriKVManager(CommonKVManager): AUX_DATA_HEADER = b"AUX_DATA" + # Implements teardown() below, so runtime PD role switching is supported. + supports_role_switch = True # The bootstrap socket carries several message kinds, so the status message # is tagged. Mori has always shipped the failure reason with it. @@ -315,6 +317,9 @@ class MoriKVManager(CommonKVManager): self.transfer_lock = threading.Lock() self._zmq_ctx = zmq.Context() self._socket_local = threading.local() + # Set by teardown() to make worker threads exit (PoC: P<->D role switch). + self._stopped = False + self._worker_threads: List[threading.Thread] = [] self._send_aux_rdma = envs.SGLANG_MORI_SEND_AUX_RDMA.get() self._register_local_buffers() if self.disaggregation_mode == DisaggregationMode.PREFILL: @@ -325,7 +330,10 @@ class MoriKVManager(CommonKVManager): self._wait_poll_ms = envs.SGLANG_MORI_WAIT_POLL_MS.get() self._transfer_timeout_ms = envs.SGLANG_MORI_TRANSFER_TIMEOUT_MS.get() for shard, queue in enumerate(self._transfer_queues): - threading.Thread( + # Track the thread so teardown() can join it: otherwise every + # P->D->P flip that re-enters PREFILL leaks _num_shards threads + # (each parked forever in FastQueue.get()). + t = threading.Thread( target=self._transfer_worker, args=(queue,), daemon=True, @@ -333,7 +341,9 @@ class MoriKVManager(CommonKVManager): f"mori-xfer-dp{self.system_dp_rank}-" f"tp{self.attn_tp_rank}-s{shard}" ), - ).start() + ) + t.start() + self._worker_threads.append(t) self._start_bootstrap_thread() elif self.disaggregation_mode == DisaggregationMode.DECODE: self._start_decode_thread() @@ -420,6 +430,11 @@ class MoriKVManager(CommonKVManager): def _transfer_worker(self, queue: FastQueue) -> None: while True: kv_chunk = queue.get() + # teardown() pushes a None sentinel to unblock get() and stop the + # worker: FastQueue.get() blocks indefinitely, so checking _stopped + # alone can never wake a parked worker during a role switch. + if kv_chunk is None: + break try: self._process_transfer_chunk(kv_chunk) except Exception as exc: @@ -683,11 +698,13 @@ class MoriKVManager(CommonKVManager): logger.debug("Room %s marked Failed via ABORT from decode", bootstrap_room) def _start_bootstrap_thread(self) -> None: + recv = self._make_worker_recv(self.server_socket) + def bootstrap_worker(): - while True: + while not self._stopped: try: - msg = self.server_socket.recv_multipart() - if not msg: + msg = recv() + if msg is None: continue tag = msg[0] @@ -705,15 +722,23 @@ class MoriKVManager(CommonKVManager): else: self._handle_transfer_message(payload) except Exception: + if self._stopped: + break logger.exception("Bootstrap worker failed") - threading.Thread(target=bootstrap_worker, daemon=True).start() + t = threading.Thread(target=bootstrap_worker, daemon=True) + t.start() + self._worker_threads.append(t) def _start_decode_thread(self) -> None: + recv = self._make_worker_recv(self.server_socket) + def decode_worker(): - while True: + while not self._stopped: try: - msg = self.server_socket.recv_multipart() + msg = recv() + if msg is None: + continue if msg and msg[0] == MoriKVManager.AUX_DATA_HEADER: self._handle_aux_data(msg) continue @@ -732,9 +757,69 @@ class MoriKVManager(CommonKVManager): failure_reason=reason, ) except Exception: + if self._stopped: + break logger.exception("Decode status worker failed") - threading.Thread(target=decode_worker, daemon=True).start() + t = threading.Thread(target=decode_worker, daemon=True) + t.start() + self._worker_threads.append(t) + + def teardown(self) -> None: + """Stop worker threads and release transport resources so this + KVManager can be discarded during a P<->D role switch. + + The KV cache pool memory is owned by the scheduler and is NOT freed + here; only mori-side registrations / sockets / engine are released. + """ + self._stopped = True + # Transfer workers (PREFILL role) park in FastQueue.get(), which has no + # timeout; push a None sentinel per shard to wake and stop them so the + # join below returns instead of leaking the thread. + for queue in getattr(self, "_transfer_queues", []): + try: + queue.put(None) + except Exception: + logger.exception("Failed to signal mori transfer worker on teardown") + # Join workers before touching their sockets: ZMQ sockets aren't + # thread-safe, so don't close server_socket while a worker may poll it. + for t in self._worker_threads: + t.join(timeout=3.0) + self._worker_threads = [] + # Drop the queues so their buffered tasks/senders are released too. + self._transfer_queues = [] + try: + self.server_socket.close(linger=0) + except Exception: + logger.exception("Failed to close mori server_socket during teardown") + # destroy() force-closes every socket in the context (incl. per-thread + # cached PUSH sockets); plain term() would block waiting on them. + try: + self._zmq_ctx.destroy(linger=0) + except Exception: + logger.exception("Failed to destroy mori zmq context during teardown") + # Deregister RDMA memory and drop the IOEngine reference. + try: + for descs in (self.kv_mem_descs, self.aux_mem_descs): + for desc in descs: + try: + self.engine.deregister_memory(desc) + except Exception: + pass + for component_descs in self.state_mem_descs: + for desc in component_descs: + try: + self.engine.deregister_memory(desc) + except Exception: + pass + finally: + self.kv_mem_descs = [] + self.aux_mem_descs = [] + self.state_mem_descs = [] + self.engine = None + logger.info( + "MoriKVManager torn down (was role=%s)", self.disaggregation_mode.value + ) def _add_remote_peer(self, register_info: KVArgsRegisterInfo) -> None: engine_key = register_info.engine_key diff --git a/python/sglang/srt/disaggregation/role_switch.py b/python/sglang/srt/disaggregation/role_switch.py new file mode 100644 index 000000000..f0b580584 --- /dev/null +++ b/python/sglang/srt/disaggregation/role_switch.py @@ -0,0 +1,284 @@ +"""Runtime prefill<->decode role switching for PD disaggregation. + +The token KV pool is role-independent and never reallocated; only the +role-specific disaggregation structures are torn down and rebuilt on a flip. +Kept out of scheduler.py to avoid growing it further. +""" + +from __future__ import annotations + +import logging +from typing import TYPE_CHECKING, Callable, Optional, Tuple + +from sglang.srt.disaggregation.common.conn import CommonKVReceiver +from sglang.srt.disaggregation.utils import DisaggregationMode +from sglang.srt.managers.io_struct import PdRoleSwitchReqInput, PdRoleSwitchReqOutput +from sglang.srt.runtime_context import get_context, get_disagg +from sglang.srt.utils import get_available_gpu_memory + +if TYPE_CHECKING: + from sglang.srt.managers.scheduler import Scheduler + +logger = logging.getLogger(__name__) + + +class PdRoleSwitchRestart(Exception): + """Break out of the current role's event loop after a successful switch.""" + + +def run_event_loop_supervisor( + scheduler: Scheduler, dispatch_once: Callable[[Scheduler], None] +) -> None: + """Re-dispatch the scheduler event loop after each runtime role switch.""" + while True: + try: + return dispatch_once(scheduler) + except PdRoleSwitchRestart: + logger.info( + "Re-dispatching event loop after PD role switch -> %s", + scheduler.disaggregation_mode.value, + ) + + +def handle_pd_role_switch( + scheduler: Scheduler, recv_req: PdRoleSwitchReqInput +) -> PdRoleSwitchReqOutput: + """Flip the scheduler's disaggregation role at runtime. The instance must be + idle; rebuild failure is fatal to the instance (no in-place rollback).""" + old_role = scheduler.disaggregation_mode.value + new_role = (recv_req.new_role or "").lower() + + def _fail(msg: str, safe_to_restore: bool = False) -> PdRoleSwitchReqOutput: + logger.warning( + "PD role switch rejected (%s -> %s): %s", old_role, new_role, msg + ) + return PdRoleSwitchReqOutput( + success=False, + message=msg, + old_role=old_role, + new_role=new_role, + safe_to_restore=safe_to_restore, + ) + + rejection = _reject_reason(scheduler, new_role) + if rejection is not None: + return _fail(*rejection) + if new_role == old_role: + return PdRoleSwitchReqOutput( + success=True, + message="already in target role", + old_role=old_role, + new_role=new_role, + ) + if not scheduler.is_fully_idle(): + return _fail( + "instance is not idle; drain all requests before switching", + safe_to_restore=True, + ) + + required_graph_gb = recv_req.decode_cuda_graph_memory_gb + # Same condition ensure_decode_cuda_graphs skips on, so the check cannot be + # bypassed while the capture still runs. + will_capture_graphs = ( + new_role == "decode" and not scheduler.tp_worker.get_decode_cuda_graph_bs() + ) + if will_capture_graphs and required_graph_gb is None: + return _fail( + "decode_cuda_graph_memory_gb is required before capturing decode graphs", + safe_to_restore=True, + ) + if will_capture_graphs and required_graph_gb is not None: + if required_graph_gb < 0: + return _fail( + "decode_cuda_graph_memory_gb must be non-negative", + safe_to_restore=True, + ) + try: + available_graph_gb = get_available_gpu_memory( + scheduler.device, scheduler.ps.gpu_id + ) + except Exception as e: + return _fail( + f"failed to check decode CUDA graph headroom: {e}", + safe_to_restore=True, + ) + if available_graph_gb < required_graph_gb: + return _fail( + "insufficient decode CUDA graph headroom: " + f"required={required_graph_gb:.3f} GB, " + f"available={available_graph_gb:.3f} GB", + safe_to_restore=True, + ) + + scheduler._pd_role_switch_in_progress = True + try: + # Teardown + role flip + rebuild are one logical atomic step. If any of + # them raises, the instance is left half-torn-down (old role released, + # new role not up) and isn't safe to serve, so mark it unhealthy. There + # is no in-place rollback. + try: + teardown_disaggregation(scheduler) + get_context().override("role_switch.flip", disaggregation_mode=new_role) + scheduler.init_disaggregation() + scheduler._sync_disaggregation_mode_to_subcomponents() + except Exception as e: + scheduler._pd_role_switch_unhealthy = True + logger.critical( + "PD role switch (%s -> %s) failed during teardown/rebuild; " + "instance unhealthy: %s", + old_role, + new_role, + e, + ) + return _fail( + f"role switch failed; instance unhealthy, restart required: {e}" + ) + + if new_role == "decode": + # Best-effort deferred capture; a failure only degrades to eager. + try: + scheduler.tp_worker.ensure_decode_cuda_graphs( + recv_req.decode_cuda_graph_bs + ) + except Exception: + logger.exception("Decode CUDA graph capture on role switch failed") + + # Break out of the old-role event loop so the supervisor re-dispatches. + scheduler._event_loop_should_restart = True + logger.info("PD role switch succeeded: %s -> %s", old_role, new_role) + return PdRoleSwitchReqOutput( + success=True, message="ok", old_role=old_role, new_role=new_role + ) + except Exception as e: + logger.exception("PD role switch failed") + return _fail(f"role switch raised: {e}") + finally: + scheduler._pd_role_switch_in_progress = False + + +def _reject_reason(scheduler: Scheduler, new_role: str) -> Optional[Tuple[str, bool]]: + """Why the switch must be rejected before draining, or None to proceed. + + Table-driven: the first failing precondition's message is returned. + """ + sa = scheduler.server_args + km = _current_kv_manager(scheduler) + # (failed?, safe to restore routing?, lazy message) + checks = ( + ( + not sa.enable_pd_role_switch, + True, + lambda: "--enable-pd-role-switch is not set on this instance", + ), + ( + scheduler._pd_role_switch_unhealthy, + False, + lambda: ( + "instance is unhealthy after a failed role switch; restart required" + ), + ), + ( + scheduler._pd_role_switch_in_progress, + False, + lambda: "another role switch is already in progress", + ), + ( + new_role not in ("prefill", "decode"), + True, + lambda: f"invalid new_role={new_role!r}", + ), + ( + scheduler.disaggregation_mode == DisaggregationMode.NULL, + True, + lambda: "instance is not running in PD disaggregation mode", + ), + ( + km is not None and not km.supports_role_switch, + True, + lambda: ( + f"transfer backend {get_disagg().disaggregation_transfer_backend!r} " + "does not support runtime role switch" + ), + ), + ( + getattr(km, "enable_staging", False), + True, + lambda: ( + "staging buffer (SGLANG_DISAGG_STAGING_BUFFER) is not " + "supported with runtime role switch" + ), + ), + ) + return next( + ((msg(), safe_to_restore) for failed, safe_to_restore, msg in checks if failed), + None, + ) + + +def _current_kv_manager(scheduler: Scheduler): + """The KV manager of the current role's disaggregation queue, or None.""" + if scheduler.disaggregation_mode == DisaggregationMode.PREFILL: + q = getattr(scheduler, "disagg_prefill_bootstrap_queue", None) + elif scheduler.disaggregation_mode == DisaggregationMode.DECODE: + q = getattr(scheduler, "disagg_decode_prealloc_queue", None) + else: + q = None + return getattr(q, "kv_manager", None) if q is not None else None + + +def teardown_disaggregation(scheduler: Scheduler) -> None: + """Release the current role's disaggregation structures (queues, metadata + buffers, KV transfer manager) so the other role can be rebuilt.""" + mode = scheduler.disaggregation_mode + if mode == DisaggregationMode.PREFILL: + q = getattr(scheduler, "disagg_prefill_bootstrap_queue", None) + if q is not None: + km = getattr(q, "kv_manager", None) + if km is not None: + km.teardown() + scheduler.disagg_prefill_bootstrap_queue = None + scheduler.disagg_prefill_inflight_queue = [] + elif mode == DisaggregationMode.DECODE: + q = getattr(scheduler, "disagg_decode_prealloc_queue", None) + if q is not None: + km = getattr(q, "kv_manager", None) + if km is not None: + km.teardown() + scheduler.disagg_decode_prealloc_queue = None + scheduler.disagg_decode_transfer_queue = None + # clear socket ctx in CommonKVReceiver + CommonKVReceiver.close_all_sockets() + scheduler.disagg_metadata_buffers = None + scheduler.req_to_metadata_buffer_idx_allocator = None + _release_prefix_cache_for_role_switch(scheduler) + + +def _release_prefix_cache_for_role_switch(scheduler: Scheduler) -> None: + """Release the prefix (radix/hicache) cache so a flip works with radix ON. + + With radix disabled (ChunkCache) the flip needs nothing here: ChunkCache + keeps no persistent prefixes and, since the instance is idle before the + switch, the allocator is already empty. This is the historical + ``--disable-radix-cache`` path, left untouched by the guard below. + + With radix (or hicache) enabled, finished prefixes stay in the tree and keep + their KV-pool slots *locked* even while idle. Carried across a role switch + that means (a) the new role would match against stale prefixes whose KV no + longer means what it did (corruption) and (b) those locked slots would leak + on every flip. Reset mirrors ``Scheduler.flush_cache``'s cache-release block + (the instance is already fully idle, checked before teardown) and, for + hicache, best-effort clears the storage backend so it is released completely. + """ + if scheduler.disable_radix_cache: + return + tree_cache = scheduler.tree_cache + if tree_cache is not None: + clear_storage = getattr(tree_cache, "clear_storage_backend", None) + if callable(clear_storage): + try: + clear_storage() + except Exception: + logger.exception("hicache storage release on role switch failed") + tree_cache.reset() + scheduler.req_to_token_pool.clear() + scheduler.token_to_kv_pool_allocator.clear() diff --git a/python/sglang/srt/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py index 361e240e6..3b4eb790b 100644 --- a/python/sglang/srt/entrypoints/engine.py +++ b/python/sglang/srt/entrypoints/engine.py @@ -1407,6 +1407,7 @@ class Engine(EngineScoreMixin, EngineBase): "load_format": tm.config_value("load_format"), "reasoning_parser": tm.config_value("reasoning_parser"), "tool_call_parser": tm.config_value("tool_call_parser"), + "disaggregation_mode": tm.config_value("disaggregation_mode"), } def init_weights_update_group( diff --git a/python/sglang/srt/entrypoints/grpc_bridge.py b/python/sglang/srt/entrypoints/grpc_bridge.py index 3b7f3b0b9..f0247abf2 100644 --- a/python/sglang/srt/entrypoints/grpc_bridge.py +++ b/python/sglang/srt/entrypoints/grpc_bridge.py @@ -411,6 +411,9 @@ class RuntimeHandle: "load_format": self.tokenizer_manager.config_value("load_format"), "reasoning_parser": self.tokenizer_manager.config_value("reasoning_parser"), "tool_call_parser": self.tokenizer_manager.config_value("tool_call_parser"), + "disaggregation_mode": self.tokenizer_manager.config_value( + "disaggregation_mode" + ), "model_type": getattr(model_config.hf_config, "model_type", None), "architectures": getattr(model_config.hf_config, "architectures", None), } diff --git a/python/sglang/srt/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py index bbe5ee339..16b5f4b5c 100644 --- a/python/sglang/srt/entrypoints/http_server.py +++ b/python/sglang/srt/entrypoints/http_server.py @@ -133,6 +133,7 @@ from sglang.srt.managers.io_struct import ( OpenSessionReqInput, ParseFunctionCallReq, PauseGenerationReqInput, + PdRoleSwitchReqInput, ProfileReq, ReleaseMemoryOccupationReqInput, ResumeMemoryOccupationReqInput, @@ -779,6 +780,9 @@ async def model_info(): "tool_call_parser": _global_state.tokenizer_manager.config_value( "tool_call_parser" ), + "disaggregation_mode": _global_state.tokenizer_manager.config_value( + "disaggregation_mode" + ), "has_image_understanding": model_config.is_image_understandable_model, "has_audio_understanding": model_config.is_audio_understandable_model, "model_type": getattr(model_config.hf_config, "model_type", None), @@ -1561,6 +1565,23 @@ async def slow_down(obj: Annotated[SlowDownReqInput, Body()], request: Request): return _create_error_response(e) +@app.api_route("/pd_role_switch", methods=["POST"]) +@auth_level(AuthLevel.ADMIN_OPTIONAL) +async def pd_role_switch( + obj: Annotated[PdRoleSwitchReqInput, Body()], request: Request +): + """Switch this instance's PD disaggregation role (prefill<->decode) at runtime. + Requires --enable-pd-role-switch; the instance must be idle.""" + try: + result = await _global_state.tokenizer_manager.pd_role_switch(obj, request) + except Exception as e: + return _create_error_response(e) + return ORJSONResponse( + msgspec_to_builtins(result), + status_code=HTTPStatus.OK if result.success else HTTPStatus.BAD_REQUEST, + ) + + @app.api_route("/load_lora_adapter", methods=["POST"]) @auth_level(AuthLevel.ADMIN_OPTIONAL) async def load_lora_adapter( diff --git a/python/sglang/srt/managers/disagg_service.py b/python/sglang/srt/managers/disagg_service.py index 3145137f6..6b5e37082 100644 --- a/python/sglang/srt/managers/disagg_service.py +++ b/python/sglang/srt/managers/disagg_service.py @@ -1,5 +1,6 @@ """Start bootstrap/kv-store-related server""" +import logging import os from sglang.srt.disaggregation.utils import ( @@ -14,14 +15,30 @@ from sglang.srt.runtime_context import ( get_serving, ) +logger = logging.getLogger(__name__) + def start_disagg_service(): # Start kv bootstrap server on prefill disagg_mode = DisaggregationMode(get_disagg().disaggregation_mode) transfer_backend = TransferBackend(get_disagg().disaggregation_transfer_backend) - if disagg_mode == DisaggregationMode.PREFILL: - # only start bootstrap server on prefill tm + # With role switching, run bootstrap on every instance (not just prefill) so + # one flipped to prefill already has it; it isn't rebuilt on flip. + start_bootstrap = disagg_mode == DisaggregationMode.PREFILL or ( + get_disagg().enable_pd_role_switch and disagg_mode != DisaggregationMode.NULL + ) + + if start_bootstrap and get_disagg().enable_pd_role_switch: + logger.warning( + "Role switch starts a bootstrap server on this instance at %s:%d. " + "If another PD instance runs on the same host, give each one a " + "distinct --disaggregation-bootstrap-port or the bind will conflict.", + get_serving().host, + get_disagg().disaggregation_bootstrap_port, + ) + + if start_bootstrap: kv_bootstrap_server_class = get_kv_class( transfer_backend, KVClassType.BOOTSTRAP_SERVER ) diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index d3f13eba8..08489c172 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -2049,6 +2049,24 @@ class SlowDownReqOutput(BaseReq, kw_only=True): pass +class PdRoleSwitchReqInput(BaseReq, kw_only=True): + # Target role; "" is an invalid sentinel rejected by the handler. + new_role: Literal["prefill", "decode", ""] = "" + # Optional decode bs to capture on a flip to decode (capture-to-fit); + # None uses the server's configured decode bs list. + decode_cuda_graph_bs: Optional[List[int]] = None + # Measured graph footprint from a matching decode peer. + decode_cuda_graph_memory_gb: Optional[float] = None + + +class PdRoleSwitchReqOutput(BaseReq, kw_only=True): + success: bool = False + message: str = "" + old_role: str = "" + new_role: str = "" + safe_to_restore: bool = False + + class AbortReq(BaseReq, kw_only=True): # Whether to abort all requests abort_all: bool = False diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index 05a221102..f154d741a 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -2225,6 +2225,7 @@ def release_req( # Callers that will recompute the KV instead (PD true-retraction rebootstrap) # pass offload_kv=False to skip the wasteful device->host copy. backup_saved = True + # The config bag reflects role flips; server_args keeps the launch role. if get_disagg().disaggregation_mode == "decode" and offload_kv: backup_saved = retraction_backup( req, diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index b574b30b2..7f8726896 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -74,6 +74,7 @@ from sglang.srt.configs.model_config import ( ) from sglang.srt.constrained.grammar_manager import GrammarManager from sglang.srt.debug_utils.pr_fix_toggle import maybe_revert_pr_fix +from sglang.srt.disaggregation import role_switch from sglang.srt.disaggregation.checksum import KvChecksumComputer from sglang.srt.disaggregation.decode import ( DecodePreallocQueue, @@ -159,6 +160,7 @@ from sglang.srt.managers.io_struct import ( MMInputsProcessError, OpenSessionReqInput, PauseGenerationReqInput, + PdRoleSwitchReqInput, ProfileReq, ReleaseMemoryOccupationReqInput, RemoveExternalCorpusReqInput, @@ -1284,6 +1286,13 @@ class Scheduler( self.hisparse_coordinator.set_decode_producer_stream(self.forward_stream) def init_running_status(self): + # Set by a runtime PD role switch to break out of the current event loop. + self._event_loop_should_restart = False + # Guards against concurrent/re-entrant PD role switches. + self._pd_role_switch_in_progress = False + # Set if a role switch tore down the old role but failed to rebuild + # either the new or the old role; the instance can no longer serve. + self._pd_role_switch_unhealthy = False # Set by the ShutdownReq handler to break the event loop for graceful shutdown. self.gracefully_exit = False self.waiting_queue: List[Req] = [] @@ -1804,6 +1813,7 @@ class Scheduler( self.weight_updater.check_weights, ), (SlowDownReqInput, self.slow_down), + (PdRoleSwitchReqInput, self.handle_pd_role_switch), ( ProfileReq, lambda req: self.profiler_manager._profile(req), @@ -2133,6 +2143,13 @@ class Scheduler( if self.external_corpus_manager is not None: self.external_corpus_manager.check_pending_load() + # A runtime PD role switch rebuilt the disaggregation structures for a new + # role. The response has already been sent above; now break out of the + # current (old-role) event loop so the supervisor can re-dispatch. + if get_disagg().enable_pd_role_switch and self._event_loop_should_restart: + self._event_loop_should_restart = False + raise role_switch.PdRoleSwitchRestart() + @staticmethod def _tokenized_requests(recv_req): if isinstance( @@ -5151,7 +5168,7 @@ class Scheduler( draft_graph_memory_usage = ( None if self.draft_worker is None else self.draft_worker.graph_memory_usage ) - ret["memory_usage"] = build_memory_usage( + memory_usage = build_memory_usage( weight_gb=self.tp_worker.model_runner.weight_load_mem_usage, kv_cache_gb=self.token_to_kv_pool_allocator.get_kvcache().mem_usage, startup_available_gb=self.startup_available_gpu_memory_gb, @@ -5160,8 +5177,29 @@ class Scheduler( target_graph_memory_usage=self.tp_worker.graph_memory_usage, draft_graph_memory_usage=draft_graph_memory_usage, ) + ret["memory_usage"] = memory_usage ret["startup_time"] = self.startup_time ret["effective_max_running_requests_per_dp"] = self.max_running_requests + # PD role switch: report this instance's role and the decode CUDA graph + # batch sizes it captured, which a router feeds back as + # PdRoleSwitchReqInput.decode_cuda_graph_bs. Unset until + # init_disaggregation runs, which also re-derives it on every flip. + disaggregation_mode = getattr(self, "disaggregation_mode", None) + if disaggregation_mode is not None: + ret["disaggregation_mode"] = disaggregation_mode.value + ret["decode_cuda_graph_bs"] = self.tp_worker.get_decode_cuda_graph_bs() + ret["decode_cuda_graph_memory_gb"] = round( + sum( + memory_usage["graph"][phase] + for phase in ( + "decode", + "target_verify", + "draft_decode", + "draft_extend", + ) + ), + 3, + ) if get_exec().moe.elastic_ep_backend is not None: from sglang.srt.elastic_ep.elastic_ep import ElasticEPStateManager @@ -5739,6 +5777,24 @@ class Scheduler( self.forward_sleep_time = t return SlowDownReqOutput() + def handle_pd_role_switch(self, recv_req: PdRoleSwitchReqInput): + return role_switch.handle_pd_role_switch(self, recv_req) + + def _sync_disaggregation_mode_to_subcomponents(self): + # Push the (possibly flipped) mode into sub-components that cache it. + # object.__setattr__ because some are frozen dataclasses. + for name in ( + "invariant_checker", + "load_inquirer", + "output_streamer", + "batch_result_processor", + ): + comp = getattr(self, name, None) + if comp is not None and hasattr(comp, "disaggregation_mode"): + object.__setattr__( + comp, "disaggregation_mode", self.disaggregation_mode + ) + def expert_distribution_handle(self, recv_req: ExpertDistributionReq): action = recv_req.action if action == ExpertDistributionReqType.START_RECORD: @@ -5818,6 +5874,15 @@ class Scheduler( def dispatch_event_loop(scheduler: Scheduler): + if scheduler.server_args.enable_pd_role_switch: + return role_switch.run_event_loop_supervisor( + scheduler, + _dispatch_event_loop_once, + ) + return _dispatch_event_loop_once(scheduler) + + +def _dispatch_event_loop_once(scheduler: Scheduler): # The live PP property asserts before torch.distributed init (MLX stub). disaggregation_mode: DisaggregationMode = scheduler.disaggregation_mode if disaggregation_mode == DisaggregationMode.NULL: diff --git a/python/sglang/srt/managers/scheduler_components/invariant_checker.py b/python/sglang/srt/managers/scheduler_components/invariant_checker.py index 6e75db3f1..16a7704a3 100644 --- a/python/sglang/srt/managers/scheduler_components/invariant_checker.py +++ b/python/sglang/srt/managers/scheduler_components/invariant_checker.py @@ -424,12 +424,13 @@ class SchedulerInvariantChecker: ) def _check_req_pool(self): - if self.disaggregation_mode == DisaggregationMode.DECODE: - req_total_size = ( - self.req_to_token_pool.size + self.req_to_token_pool.pre_alloc_size - ) - else: - req_total_size = self.req_to_token_pool.size + # Account for pre-alloc headroom whenever the pool has it. The decode + # pool always does; with runtime P<->D role switching a prefill instance + # may also hold a decode-flavored (pre-alloc) pool, so key off the pool + # itself rather than the current role. + req_total_size = ( + self.req_to_token_pool.size + self.req_to_token_pool.pre_alloc_size + ) session_req_count = self.pool_stats_observer.session_held_req_count() if len(self.req_to_token_pool.free_slots) + session_req_count != req_total_size: diff --git a/python/sglang/srt/managers/tokenizer_control_mixin.py b/python/sglang/srt/managers/tokenizer_control_mixin.py index 69ef86803..0170ac34e 100644 --- a/python/sglang/srt/managers/tokenizer_control_mixin.py +++ b/python/sglang/srt/managers/tokenizer_control_mixin.py @@ -9,6 +9,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple import fastapi +from sglang.srt.disaggregation.utils import DisaggregationMode from sglang.srt.managers.communicator import FanOutCommunicator from sglang.srt.managers.io_struct import ( AddExternalCorpusReqInput, @@ -48,6 +49,8 @@ from sglang.srt.managers.io_struct import ( LoadLoRAAdapterReqOutput, LoRAUpdateOutput, OpenSessionReqInput, + PdRoleSwitchReqInput, + PdRoleSwitchReqOutput, ProfileReq, ProfileReqOutput, ProfileReqType, @@ -77,6 +80,7 @@ from sglang.srt.managers.io_struct import ( ) from sglang.srt.managers.load_snapshot import LoadSnapshot from sglang.srt.runtime_context import ( + get_disagg, get_lora, get_parallel, get_serving, @@ -115,6 +119,7 @@ _COMMUNICATOR_SPECS = [ ("resume_memory_occupation", ResumeMemoryOccupationReqOutput), ("check_weights", CheckWeightsReqOutput), ("slow_down", SlowDownReqOutput), + ("pd_role_switch", PdRoleSwitchReqOutput), ("flush_cache", FlushCacheReqOutput), ("add_external_corpus", AddExternalCorpusReqOutput), ("remove_external_corpus", RemoveExternalCorpusReqOutput), @@ -840,6 +845,42 @@ class TokenizerControlMixin: self.auto_create_handle_loop() await self.slow_down_communicator(obj) + async def pd_role_switch( + self: TokenizerManager, + obj: PdRoleSwitchReqInput, + request: Optional[fastapi.Request] = None, + ) -> PdRoleSwitchReqOutput: + self.auto_create_handle_loop() + if not self.server_args.enable_pd_role_switch: + return PdRoleSwitchReqOutput( + success=False, + message="--enable-pd-role-switch is not set on this server", + old_role=get_disagg().disaggregation_mode, + new_role=obj.new_role, + safe_to_restore=True, + ) + results = await self.pd_role_switch_communicator(obj) + all_success = all(r.success for r in results) + safe_to_restore = bool(results) and all(r.safe_to_restore for r in results) + if all_success: + # Keep the tokenizer-manager's view of the role in sync so future + # control ops and bootstrap routing behave consistently. + self.record_config_updates( + "tokenizer.pd_role_switch", disaggregation_mode=obj.new_role + ) + self.disaggregation_mode = DisaggregationMode(obj.new_role) + msg = "ok" + else: + # Surface only the failing workers' messages. + msg = "; ".join(r.message for r in results if not r.success) + return PdRoleSwitchReqOutput( + success=all_success, + message=msg, + old_role=results[0].old_role if results else "", + new_role=obj.new_role, + safe_to_restore=safe_to_restore, + ) + async def get_internal_state(self: TokenizerManager) -> List[Dict[Any, Any]]: self.auto_create_handle_loop() req = GetInternalStateReq() diff --git a/python/sglang/srt/managers/tp_worker.py b/python/sglang/srt/managers/tp_worker.py index c05f6c621..a9067ac1f 100644 --- a/python/sglang/srt/managers/tp_worker.py +++ b/python/sglang/srt/managers/tp_worker.py @@ -471,6 +471,18 @@ class TpModelWorker(BaseTpWorker): for mr in self.model_runner_list[1:]: mr.init_cuda_graphs(capture_decode_cuda_graph=capture_decode_cuda_graph) + def ensure_decode_cuda_graphs(self, capture_bs: Optional[List[int]] = None): + """Idempotently capture decode cuda graphs for all model runners (used + for the on-flip capture during a runtime PD role switch).""" + self.model_runner.ensure_decode_cuda_graphs(capture_bs) + for mr in self.model_runner_list[1:]: + mr.ensure_decode_cuda_graphs(capture_bs) + + def get_decode_cuda_graph_bs(self) -> List[int]: + """Decode bs captured as CUDA graphs (empty on a not-yet-flipped prefill, + or on a runner that never allocates a KV pool, e.g. the MLX stub).""" + return list(getattr(self.model_runner, "decode_cuda_graph_capture_bs", [])) + def start_startup_weight_load(self) -> None: """Start deferred checkpoint prefetching for all model runners.""" self.model_runner.start_startup_weight_load() diff --git a/python/sglang/srt/mem_cache/kv_cache_configurator.py b/python/sglang/srt/mem_cache/kv_cache_configurator.py index f1d761c45..d6a28b2f9 100644 --- a/python/sglang/srt/mem_cache/kv_cache_configurator.py +++ b/python/sglang/srt/mem_cache/kv_cache_configurator.py @@ -1009,32 +1009,34 @@ class KVCacheConfigurator: def _build_req_to_token_pool(self, *, max_num_reqs: int) -> ReqToTokenPool: extra_max_context_len = get_req_to_token_extra_context_len() - if get_disagg().disaggregation_mode == "decode": - # Extra slots for pre-allocated requests - pre_alloc_size = get_disagg().disaggregation_decode_extra_slots + disagg = get_disagg() + if disagg.disaggregation_mode == "decode" or disagg.enable_pd_role_switch: + # A flip-capable prefill needs the decode pool shape, and the extra-slot + # default is only computed for a decode launch. + pre_alloc_size = disagg.disaggregation_decode_extra_slots + if disagg.enable_pd_role_switch: + pre_alloc_size = pre_alloc_size or 0 if self.mambaish_config: - req_to_token_pool = self._build_hybrid_mamba_decode_req_pool( + return self._build_hybrid_mamba_decode_req_pool( max_num_reqs=max_num_reqs, extra_max_context_len=extra_max_context_len, pre_alloc_size=pre_alloc_size, ) - else: - req_to_token_pool = self._build_decode_req_pool( - max_num_reqs=max_num_reqs, - extra_max_context_len=extra_max_context_len, - pre_alloc_size=pre_alloc_size, - ) - elif self.mambaish_config: - req_to_token_pool = self._build_hybrid_req_pool( + return self._build_decode_req_pool( + max_num_reqs=max_num_reqs, + extra_max_context_len=extra_max_context_len, + pre_alloc_size=pre_alloc_size, + ) + + if self.mambaish_config: + return self._build_hybrid_req_pool( max_num_reqs=max_num_reqs, extra_max_context_len=extra_max_context_len, ) - else: - req_to_token_pool = self._build_default_req_pool( - max_num_reqs=max_num_reqs, - extra_max_context_len=extra_max_context_len, - ) - return req_to_token_pool + return self._build_default_req_pool( + max_num_reqs=max_num_reqs, + extra_max_context_len=extra_max_context_len, + ) def _get_mamba_layer_ids_for_req_pool(self) -> list: mamba_layer_ids = [ diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py index b18e8964e..65604894e 100644 --- a/python/sglang/srt/mem_cache/memory_pool.py +++ b/python/sglang/srt/mem_cache/memory_pool.py @@ -273,6 +273,10 @@ class ReqToTokenPool: """A memory pool that maps a request to its token locations.""" enable_mamba_extra_buffer_lazy: bool = False + # Extra pre-allocation headroom (reserved for in-transfer decode requests). + # 0 for a plain pool; the decode-flavored pool (DecodeReqToTokenPool) sets a + # positive value. Declared here so callers can read it without getattr. + pre_alloc_size: int = 0 # Class default: some decode pools borrow another __init__ (see # DecodeReqToTokenPool) but inherit alloc_rows. _on_alloc_rows: Optional[Callable[[List[int]], None]] = None diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 7b0fdf83a..87dc0bfb0 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -92,6 +92,7 @@ from sglang.srt.mem_cache.kv_cache_configurator import ( from sglang.srt.mem_cache.kv_index_translator import KVIndexTranslator from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool from sglang.srt.model_executor.cuda_graph_config import ( + Backend, cuda_graph_fully_disabled, ) from sglang.srt.model_executor.forward_batch_info import ( @@ -107,6 +108,7 @@ from sglang.srt.model_executor.graph_memory_usage import ( replace_graph_memory_usage, replace_graph_time_usage, ) +from sglang.srt.model_executor.graph_shared_output import GraphSharedOutput from sglang.srt.model_executor.model_runner_components import misc_utils from sglang.srt.model_executor.model_runner_components.attention_backend_setup import ( build_attention_backends, @@ -936,6 +938,11 @@ class ModelRunner: self.init_indexer_capturer() self.graph_shared_output = None + # Set once real decode CUDA graphs are captured (makes on-flip role-switch + # capture idempotent). + self.decode_cuda_graph_captured = False + # Captured decode bs; exposed via /get_server_info for role-switch queries. + self.decode_cuda_graph_capture_bs: list[int] = [] def maybe_init_hisparse_coordinator(self): if not self.enable_hisparse: @@ -1492,6 +1499,51 @@ class ModelRunner: capture.time_usage, phases=("decode", "target_verify", "draft_decode"), ) + # Bookkeeping for the PD role switch: mark the graphs as captured (makes + # the on-flip capture idempotent) and record the captured bs so it can be + # queried via /get_server_info. + self.decode_cuda_graph_captured = self.decode_cuda_graph_runner is not None + self.decode_cuda_graph_capture_bs = list( + getattr(self.decode_cuda_graph_runner, "capture_bs", []) or [] + ) + + def ensure_decode_cuda_graphs(self, capture_bs: Optional[list[int]] = None): + """Idempotently capture decode CUDA graphs after startup. + + Used by the PD role switch: an instance launched as prefill runs fully + eager (decode CUDA graph disabled). On the first flip to decode we + enable the decode CUDA graph and capture it here, so the flipped + instance replays decode graphs instead of running eager. + """ + if self.decode_cuda_graph_captured: + logger.info("Decode CUDA graphs already captured; skipping re-capture.") + return + + cfg = get_exec().graph.cuda_graph_config + was_disabled = cfg is not None and cfg.decode.backend == Backend.DISABLED + if was_disabled: + # Prefill was launched with the decode CUDA graph disabled; enable it + # for the decode role. + logger.info( + "Enabling decode CUDA graph on role switch (was disabled at startup)." + ) + cfg.decode.backend = Backend.FULL + get_context().override( + "model_runner.ensure_decode_cuda_graphs", disable_cuda_graph=False + ) + + if capture_bs: + # Capture-to-fit: only the requested (router-sized) batch sizes. + filtered_bs = sorted({int(b) for b in capture_bs if int(b) > 0}) + if filtered_bs: + cfg.decode.bs = filtered_bs + + if was_disabled: + # graph_shared_output is skipped at startup when decode is disabled, + # so build it now (before the decode runner reads its logits buffer). + self.graph_shared_output = GraphSharedOutput.create_for_model_runner(self) + + self.init_decode_cuda_graph() def init_prefill_cuda_graph(self, force_for_draft_worker: bool = False): self.prefill_cuda_graph_runner = None diff --git a/python/sglang/srt/observability/metrics_collector.py b/python/sglang/srt/observability/metrics_collector.py index 79d55e574..d58f9bcf1 100644 --- a/python/sglang/srt/observability/metrics_collector.py +++ b/python/sglang/srt/observability/metrics_collector.py @@ -1124,8 +1124,11 @@ class SchedulerMetricsCollector(_StatLoggerDIMixin): ) collector: Optional[SchedulerMetricsCollector] = None if enable_metrics: - engine_type = DisaggregationMode.to_engine_type( - get_disagg().disaggregation_mode + # Keep one metric series across role flips. + engine_type = ( + "dynamic" + if get_disagg().enable_pd_role_switch + else DisaggregationMode.to_engine_type(get_disagg().disaggregation_mode) ) labels = { "model_name": get_serving().served_model_name, diff --git a/rust/sglang-server/src/api_server/common.rs b/rust/sglang-server/src/api_server/common.rs index 5e6012c79..4c64a4e9a 100644 --- a/rust/sglang-server/src/api_server/common.rs +++ b/rust/sglang-server/src/api_server/common.rs @@ -94,6 +94,7 @@ async fn model_info(State(state): State>) -> Response { // selected parser into `server_args` before the scheduler forks. "reasoning_parser": sa.reasoning_parser, "tool_call_parser": sa.tool_call_parser, + "disaggregation_mode": sa.disaggregation_mode, }); ( StatusCode::OK, diff --git a/rust/sglang-server/src/message/config.rs b/rust/sglang-server/src/message/config.rs index e8f102080..8e44dd279 100644 --- a/rust/sglang-server/src/message/config.rs +++ b/rust/sglang-server/src/message/config.rs @@ -317,7 +317,9 @@ impl<'py> pyo3::FromPyObject<'_, 'py> for PreferredSamplingParams { from_py_object, module = "sglang.srt.rust_extensions._server" )] -#[derive(Clone, Copy, Debug, PartialEq, Eq)] +// Lowercase to match the values Python reports for the same field. +#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize)] +#[serde(rename_all = "lowercase")] pub enum DisaggregationMode { /// Unified prefill + decode. Null, @@ -619,6 +621,14 @@ mod tests { assert_eq!(ServerArgs::default().bind(), "127.0.0.1:30000"); } + #[test] + fn disaggregation_mode_wire_values_match_python() { + let json = |m| serde_json::to_string(&m).unwrap(); + assert_eq!(json(DisaggregationMode::Null), "\"null\""); + assert_eq!(json(DisaggregationMode::Prefill), "\"prefill\""); + assert_eq!(json(DisaggregationMode::Decode), "\"decode\""); + } + #[test] fn pd_role_derivations() { let prefill = ServerArgs { diff --git a/sgl-model-gateway/bindings/python/src/sglang_router/mini_lb.py b/sgl-model-gateway/bindings/python/src/sglang_router/mini_lb.py index cb8feefad..3cb49fa53 100644 --- a/sgl-model-gateway/bindings/python/src/sglang_router/mini_lb.py +++ b/sgl-model-gateway/bindings/python/src/sglang_router/mini_lb.py @@ -111,6 +111,42 @@ class MiniLoadBalancer: self.decode_urls[didx], ) + def current_role_and_port(self, worker_url): + """Return (role, bootstrap_port) of a registered server, or (None, None) + if it is not in either routing list.""" + if worker_url in self.prefill_urls: + return ( + "prefill", + self.prefill_bootstrap_ports[self.prefill_urls.index(worker_url)], + ) + if worker_url in self.decode_urls: + return "decode", None + return None, None + + def remove_worker(self, worker_url): + """Drop a server from both routing lists so no new requests are sent to + it (used to quiesce it before a role switch).""" + if worker_url in self.decode_urls: + self.decode_urls.remove(worker_url) + if worker_url in self.prefill_urls: + idx = self.prefill_urls.index(worker_url) + self.prefill_urls.pop(idx) + self.prefill_bootstrap_ports.pop(idx) + + def add_worker(self, worker_url, role, bootstrap_port=None): + """Register a server under a role in the routing lists.""" + if role == "prefill": + self.prefill_urls.append(worker_url) + self.prefill_bootstrap_ports.append(bootstrap_port or 8998) + elif role == "decode": + self.decode_urls.append(worker_url) + + def apply_role_switch(self, worker_url, new_role, bootstrap_port=None): + """Move a server between the prefill and decode routing lists after its + role has been switched on the backend. Idempotent.""" + self.remove_worker(worker_url) + self.add_worker(worker_url, new_role, bootstrap_port) + async def generate( self, modified_request, prefill_server, decode_server, endpoint ) -> ORJSONResponse: @@ -255,6 +291,69 @@ async def health_generate(): return Response(status_code=200) +async def _post_role_switch(worker_url, body): + """POST the role switch to a backend server; return (status, json).""" + try: + async with aiohttp.ClientSession( + timeout=aiohttp.ClientTimeout(total=lb.timeout) + ) as session: + async with session.post(f"{worker_url}/pd_role_switch", json=body) as resp: + return resp.status, await resp.json() + except Exception as e: # transport error -> report as a failure + return 502, {"success": False, "message": str(e)} + + +@app.post("/pd_role_switch") +async def pd_role_switch(request_data: dict): + """Switch a running server's PD role (prefill<->decode) at runtime and + update the LB's routing lists. Body: {"worker_url", "new_role": + "prefill"|"decode", "bootstrap_port"?, "decode_cuda_graph_bs"?, + "decode_cuda_graph_memory_gb"?, "drain"?, "drain_timeout_secs"?}. + + The backend rejects a switch unless the instance is idle. To make this + safe while serving, by default the LB first removes the server from its + routing lists (so no new requests arrive), then retries the switch while + the server drains its in-flight requests, and only then registers it + under the new role. A failed server is restored only when the backend + confirms that no role state changed.""" + worker_url = request_data.get("worker_url") + new_role = request_data.get("new_role") + if worker_url is None: + raise HTTPException(status_code=400, detail="worker_url is required") + if new_role not in ("prefill", "decode"): + raise HTTPException(status_code=400, detail=f"invalid new_role={new_role!r}") + + drain = request_data.get("drain", True) + drain_timeout = request_data.get("drain_timeout_secs", 300) + old_role, old_port = lb.current_role_and_port(worker_url) + + body = {"new_role": new_role} + for field in ("decode_cuda_graph_bs", "decode_cuda_graph_memory_gb"): + if request_data.get(field) is not None: + body[field] = request_data[field] + + # Stop routing new requests to this server so it can drain to idle. + if drain and old_role is not None: + lb.remove_worker(worker_url) + + deadline = asyncio.get_event_loop().time() + drain_timeout + while True: + status, result = await _post_role_switch(worker_url, body) + if status == 200 and result.get("success", False): + break + # The backend rejects while not idle; keep retrying as it drains. + not_idle = "not idle" in (result.get("message", "") or "").lower() + if drain and not_idle and asyncio.get_event_loop().time() < deadline: + await asyncio.sleep(1.0) + continue + if drain and old_role is not None and result.get("safe_to_restore", False): + lb.add_worker(worker_url, old_role, old_port) + return ORJSONResponse(content=result, status_code=status) + + lb.apply_role_switch(worker_url, new_role, request_data.get("bootstrap_port")) + return ORJSONResponse(content=result, status_code=200) + + @app.post("/flush_cache") async def flush_cache(timeout: Optional[float] = None): # `timeout` must reach the workers. The scheduler treats a missing or diff --git a/sgl-model-gateway/bindings/python/tests/test_mini_lb_role_switch.py b/sgl-model-gateway/bindings/python/tests/test_mini_lb_role_switch.py new file mode 100644 index 000000000..ecef27092 --- /dev/null +++ b/sgl-model-gateway/bindings/python/tests/test_mini_lb_role_switch.py @@ -0,0 +1,84 @@ +import asyncio + +import pytest +from sglang_router import mini_lb + + +@pytest.mark.parametrize( + ("result", "restored"), + [ + ( + { + "success": False, + "message": "instance is not idle", + "safe_to_restore": True, + }, + True, + ), + ( + { + "success": False, + "message": "instance unhealthy, restart required", + "safe_to_restore": False, + }, + False, + ), + ({"success": False, "message": "connection lost"}, False), + ], +) +def test_failed_role_switch_restores_only_healthy_worker(monkeypatch, result, restored): + worker_url = "http://prefill:8000" + load_balancer = mini_lb.MiniLoadBalancer.__new__(mini_lb.MiniLoadBalancer) + load_balancer.timeout = 1 + load_balancer.prefill_urls = [worker_url] + load_balancer.prefill_bootstrap_ports = [8998] + load_balancer.decode_urls = ["http://decode:8000"] + monkeypatch.setattr(mini_lb, "lb", load_balancer) + + async def post_role_switch(*_args, **_kwargs): + return 400, result + + monkeypatch.setattr(mini_lb, "_post_role_switch", post_role_switch) + response = asyncio.run( + mini_lb.pd_role_switch( + { + "worker_url": worker_url, + "new_role": "decode", + "drain_timeout_secs": 0, + } + ) + ) + + assert response.status_code == 400 + assert (worker_url in load_balancer.prefill_urls) is restored + + +def test_role_switch_forwards_decode_graph_requirements(monkeypatch): + worker_url = "http://prefill:8000" + load_balancer = mini_lb.MiniLoadBalancer.__new__(mini_lb.MiniLoadBalancer) + load_balancer.timeout = 1 + load_balancer.prefill_urls = [worker_url] + load_balancer.prefill_bootstrap_ports = [8998] + load_balancer.decode_urls = ["http://decode:8000"] + monkeypatch.setattr(mini_lb, "lb", load_balancer) + sent_body = {} + + async def post_role_switch(_worker_url, body): + sent_body.update(body) + return 200, {"success": True, "message": "ok"} + + monkeypatch.setattr(mini_lb, "_post_role_switch", post_role_switch) + response = asyncio.run( + mini_lb.pd_role_switch( + { + "worker_url": worker_url, + "new_role": "decode", + "decode_cuda_graph_bs": [1, 2, 4], + "decode_cuda_graph_memory_gb": 1.25, + } + ) + ) + + assert response.status_code == 200 + assert sent_body["decode_cuda_graph_bs"] == [1, 2, 4] + assert sent_body["decode_cuda_graph_memory_gb"] == 1.25 diff --git a/test/registered/unit/disaggregation/test_pd_role_switch.py b/test/registered/unit/disaggregation/test_pd_role_switch.py new file mode 100644 index 000000000..617ee3d2d --- /dev/null +++ b/test/registered/unit/disaggregation/test_pd_role_switch.py @@ -0,0 +1,653 @@ +import argparse +import concurrent.futures +import unittest +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +from sglang.srt import runtime_context as rc # noqa: E402 +from sglang.srt.disaggregation import role_switch # noqa: E402 +from sglang.srt.disaggregation.utils import DisaggregationMode # noqa: E402 +from sglang.srt.managers.io_struct import ( # noqa: E402 + PdRoleSwitchReqInput, + PdRoleSwitchReqOutput, +) +from sglang.srt.managers.scheduler import Scheduler # noqa: E402 +from sglang.srt.server_args import ServerArgs # noqa: E402 +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + + +class TestPdRoleSwitchServerArg(unittest.TestCase): + def test_cli_flag_parses(self): + parser = argparse.ArgumentParser() + ServerArgs.add_cli_args(parser) + + off = parser.parse_args(["--model-path", "dummy"]) + self.assertFalse(off.enable_pd_role_switch) + + on = parser.parse_args(["--model-path", "dummy", "--enable-pd-role-switch"]) + self.assertTrue(on.enable_pd_role_switch) + + +class TestHandlePdRoleSwitch(unittest.TestCase): + """Cover the control-plane contract of Scheduler.handle_pd_role_switch. + + Only the role-flip *decision* logic is exercised here (no GPU): the heavy + teardown/rebuild is mocked, so this asserts the guard branches and the + orchestration order without standing up a model. + """ + + def setUp(self): + rc.reset_context() + + def tearDown(self): + rc.reset_context() + + def _scheduler(self, mode, *, enable=True, idle=True): + s = Scheduler.__new__(Scheduler) + s.disaggregation_mode = mode + sa = ServerArgs( + model_path="dummy", + disaggregation_mode=mode.value, + enable_pd_role_switch=enable, + ) + rc.get_context().set_server_args(sa) + s.server_args = sa + s.is_fully_idle = MagicMock(return_value=idle) + teardown_patcher = patch.object(role_switch, "teardown_disaggregation") + s.teardown_disaggregation = teardown_patcher.start() + self.addCleanup(teardown_patcher.stop) + s.init_disaggregation = MagicMock() + s._sync_disaggregation_mode_to_subcomponents = MagicMock() + s._event_loop_should_restart = False + s._pd_role_switch_in_progress = False + s._pd_role_switch_unhealthy = False + s.tp_worker = MagicMock() + return s + + def test_rejected_when_flag_disabled(self): + s = self._scheduler(DisaggregationMode.PREFILL, enable=False) + out = Scheduler.handle_pd_role_switch( + s, PdRoleSwitchReqInput(new_role="decode") + ) + self.assertIsInstance(out, PdRoleSwitchReqOutput) + self.assertFalse(out.success) + self.assertTrue(out.safe_to_restore) + self.assertIn("enable-pd-role-switch", out.message) + s.teardown_disaggregation.assert_not_called() + + def test_rejected_on_invalid_role(self): + s = self._scheduler(DisaggregationMode.PREFILL) + out = Scheduler.handle_pd_role_switch(s, PdRoleSwitchReqInput(new_role="both")) + self.assertFalse(out.success) + self.assertTrue(out.safe_to_restore) + self.assertIn("invalid new_role", out.message) + s.teardown_disaggregation.assert_not_called() + + def test_rejected_when_not_in_pd_mode(self): + s = self._scheduler(DisaggregationMode.NULL) + out = Scheduler.handle_pd_role_switch( + s, PdRoleSwitchReqInput(new_role="decode") + ) + self.assertFalse(out.success) + self.assertTrue(out.safe_to_restore) + self.assertIn("not running in PD", out.message) + s.teardown_disaggregation.assert_not_called() + + def test_same_role_is_noop(self): + s = self._scheduler(DisaggregationMode.PREFILL) + out = Scheduler.handle_pd_role_switch( + s, PdRoleSwitchReqInput(new_role="prefill") + ) + self.assertTrue(out.success) + self.assertEqual(out.message, "already in target role") + s.teardown_disaggregation.assert_not_called() + s.init_disaggregation.assert_not_called() + self.assertFalse(s._event_loop_should_restart) + + def test_rejected_when_not_idle(self): + s = self._scheduler(DisaggregationMode.PREFILL, idle=False) + out = Scheduler.handle_pd_role_switch( + s, PdRoleSwitchReqInput(new_role="decode") + ) + self.assertFalse(out.success) + self.assertTrue(out.safe_to_restore) + self.assertIn("not idle", out.message) + s.teardown_disaggregation.assert_not_called() + + def test_rejected_when_decode_graph_headroom_is_missing(self): + s = self._scheduler(DisaggregationMode.PREFILL) + s.tp_worker.get_decode_cuda_graph_bs.return_value = [] + out = Scheduler.handle_pd_role_switch( + s, PdRoleSwitchReqInput(new_role="decode") + ) + + self.assertFalse(out.success) + self.assertTrue(out.safe_to_restore) + self.assertIn("decode_cuda_graph_memory_gb is required", out.message) + s.teardown_disaggregation.assert_not_called() + + def test_rejected_when_decode_graph_headroom_is_insufficient(self): + s = self._scheduler(DisaggregationMode.PREFILL) + s.device = "cuda" + s.ps = SimpleNamespace(gpu_id=0) + s.tp_worker.get_decode_cuda_graph_bs.return_value = [] + with patch.object(role_switch, "get_available_gpu_memory", return_value=0.5): + out = Scheduler.handle_pd_role_switch( + s, + PdRoleSwitchReqInput( + new_role="decode", + decode_cuda_graph_memory_gb=1.0, + ), + ) + + self.assertFalse(out.success) + self.assertTrue(out.safe_to_restore) + self.assertIn("insufficient decode CUDA graph headroom", out.message) + s.teardown_disaggregation.assert_not_called() + self.assertEqual(rc.get_disagg().disaggregation_mode, "prefill") + + def test_decode_graph_headroom_allows_flip(self): + s = self._scheduler(DisaggregationMode.PREFILL) + s.device = "cuda" + s.ps = SimpleNamespace(gpu_id=0) + s.tp_worker.get_decode_cuda_graph_bs.return_value = [] + with patch.object(role_switch, "get_available_gpu_memory", return_value=1.0): + out = Scheduler.handle_pd_role_switch( + s, + PdRoleSwitchReqInput( + new_role="decode", + decode_cuda_graph_memory_gb=1.0, + ), + ) + + self.assertTrue(out.success) + s.teardown_disaggregation.assert_called_once_with(s) + + def test_successful_flip_orchestration(self): + s = self._scheduler(DisaggregationMode.PREFILL) + out = Scheduler.handle_pd_role_switch( + s, PdRoleSwitchReqInput(new_role="decode") + ) + + self.assertTrue(out.success) + self.assertEqual(out.old_role, "prefill") + self.assertEqual(out.new_role, "decode") + # Orchestration: drain -> teardown -> flip config bag -> rebuild -> signal. + s.teardown_disaggregation.assert_called_once_with(s) + self.assertEqual(rc.get_disagg().disaggregation_mode, "decode") + # The pristine startup record is never mutated. + self.assertEqual(s.server_args.disaggregation_mode, "prefill") + s.init_disaggregation.assert_called_once() + s._sync_disaggregation_mode_to_subcomponents.assert_called_once() + self.assertTrue(s._event_loop_should_restart) + # Flip to decode ensures decode CUDA graphs exist (idempotent capture). + s.tp_worker.ensure_decode_cuda_graphs.assert_called_once() + # The in-progress guard is released after a successful flip. + self.assertFalse(s._pd_role_switch_in_progress) + + def test_flip_to_prefill_skips_decode_graph_capture(self): + s = self._scheduler(DisaggregationMode.DECODE) + out = Scheduler.handle_pd_role_switch( + s, PdRoleSwitchReqInput(new_role="prefill") + ) + self.assertTrue(out.success) + self.assertEqual(out.new_role, "prefill") + s.init_disaggregation.assert_called_once() + # Flipping to prefill must not capture decode graphs. + s.tp_worker.ensure_decode_cuda_graphs.assert_not_called() + self.assertTrue(s._event_loop_should_restart) + + def test_rejected_when_switch_in_progress(self): + s = self._scheduler(DisaggregationMode.PREFILL) + s._pd_role_switch_in_progress = True + out = Scheduler.handle_pd_role_switch( + s, PdRoleSwitchReqInput(new_role="decode") + ) + self.assertFalse(out.success) + self.assertFalse(out.safe_to_restore) + self.assertIn("in progress", out.message) + s.teardown_disaggregation.assert_not_called() + + def test_rejected_when_unhealthy(self): + s = self._scheduler(DisaggregationMode.PREFILL) + s._pd_role_switch_unhealthy = True + out = Scheduler.handle_pd_role_switch( + s, PdRoleSwitchReqInput(new_role="decode") + ) + self.assertFalse(out.success) + self.assertFalse(out.safe_to_restore) + self.assertIn("unhealthy", out.message) + s.teardown_disaggregation.assert_not_called() + + def test_rebuild_failure_marks_unhealthy_and_notifies(self): + s = self._scheduler(DisaggregationMode.PREFILL) + # Rebuild of the new role fails after the old role was torn down. + s.init_disaggregation = MagicMock(side_effect=RuntimeError("boom")) + + out = Scheduler.handle_pd_role_switch( + s, PdRoleSwitchReqInput(new_role="decode") + ) + + # Fail loud (notify), mark unhealthy, no in-place rollback attempt. + self.assertFalse(out.success) + self.assertFalse(out.safe_to_restore) + self.assertIn("unhealthy", out.message) + self.assertIn("restart", out.message) + self.assertTrue(s._pd_role_switch_unhealthy) + self.assertFalse(s._event_loop_should_restart) + self.assertFalse(s._pd_role_switch_in_progress) + # Teardown + rebuild attempted exactly once (no rollback). + self.assertEqual(s.teardown_disaggregation.call_count, 1) + self.assertEqual(s.init_disaggregation.call_count, 1) + s._sync_disaggregation_mode_to_subcomponents.assert_not_called() + # A subsequent switch is rejected because the instance is unhealthy. + out2 = Scheduler.handle_pd_role_switch( + s, PdRoleSwitchReqInput(new_role="prefill") + ) + self.assertFalse(out2.success) + self.assertIn("unhealthy", out2.message) + + def test_teardown_failure_marks_unhealthy(self): + """Teardown, the role flip and rebuild are one atomic step: a failure + during teardown (not only rebuild) must also mark the instance unhealthy + and must not proceed to rebuild.""" + s = self._scheduler(DisaggregationMode.PREFILL) + s.teardown_disaggregation.side_effect = RuntimeError("boom") + + out = Scheduler.handle_pd_role_switch( + s, PdRoleSwitchReqInput(new_role="decode") + ) + + self.assertFalse(out.success) + self.assertFalse(out.safe_to_restore) + self.assertIn("unhealthy", out.message) + self.assertIn("restart", out.message) + self.assertTrue(s._pd_role_switch_unhealthy) + self.assertFalse(s._event_loop_should_restart) + self.assertFalse(s._pd_role_switch_in_progress) + # Teardown raised, so rebuild is never attempted. + self.assertEqual(s.teardown_disaggregation.call_count, 1) + s.init_disaggregation.assert_not_called() + s._sync_disaggregation_mode_to_subcomponents.assert_not_called() + + +class TestPdRoleSwitchReqSerialization(unittest.TestCase): + """Guard the wire contract of the /pd_role_switch req/resp structs. + + These caught real breakages when upstream moved BaseReq to msgspec: the + request must accept an optional decode_cuda_graph_bs body field, and the + response must be encodable for the HTTP layer (msgspec_to_builtins). + """ + + def test_req_accepts_optional_decode_cuda_graph_bs(self): + req = PdRoleSwitchReqInput( + new_role="decode", + decode_cuda_graph_bs=[1, 2, 4], + decode_cuda_graph_memory_gb=1.25, + ) + self.assertEqual(req.new_role, "decode") + self.assertEqual(req.decode_cuda_graph_bs, [1, 2, 4]) + self.assertEqual(req.decode_cuda_graph_memory_gb, 1.25) + # Field is optional and defaults to None. + default_req = PdRoleSwitchReqInput(new_role="prefill") + self.assertIsNone(default_req.decode_cuda_graph_bs) + self.assertIsNone(default_req.decode_cuda_graph_memory_gb) + + def test_resp_is_json_encodable(self): + from sglang.srt.utils.msgspec_utils import msgspec_to_builtins + + out = PdRoleSwitchReqOutput( + success=True, message="ok", old_role="prefill", new_role="decode" + ) + d = msgspec_to_builtins(out) + self.assertEqual(d["success"], True) + self.assertEqual(d["old_role"], "prefill") + self.assertEqual(d["new_role"], "decode") + self.assertEqual(d["message"], "ok") + self.assertEqual(d["safe_to_restore"], False) + + +class TestPdRoleSwitchStartupValidation(unittest.TestCase): + """--enable-pd-role-switch only rebuilds the small role-specific disagg + structures on a flip; the per-role buffers of DP attention / EP / MoE + all-to-all / pipeline parallelism are sized at startup and not rebuilt, so + a flip with those on would silently deadlock. The PD arg hook must reject + the combination up-front instead of failing at flip time.""" + + def _sa(self, **kw): + base = dict( + disaggregation_transfer_backend="mori", + disaggregation_mode="prefill", + enable_pd_role_switch=True, + enable_dp_attention=False, + ep_size=1, + moe_a2a_backend="none", + pp_size=1, + dp_size=1, + dcp_size=1, + speculative_algorithm=None, + ) + base.update(kw) + return SimpleNamespace(**base) + + def _run(self, sa): + from sglang.srt.arg_groups.pd_disaggregation_hook import ( + handle_pd_disaggregation, + ) + + handle_pd_disaggregation(sa) + + def test_pure_tp_role_switch_accepted(self): + # No raise for the validated pure-TP configuration. + self._run(self._sa()) + + def test_reject_dp_attention(self): + with self.assertRaises(ValueError) as ctx: + self._run(self._sa(enable_dp_attention=True)) + self.assertIn("DP attention", str(ctx.exception)) + + def test_reject_expert_parallelism(self): + with self.assertRaises(ValueError) as ctx: + self._run(self._sa(ep_size=8)) + self.assertIn("expert parallelism", str(ctx.exception)) + + def test_reject_moe_a2a(self): + with self.assertRaises(ValueError) as ctx: + self._run(self._sa(moe_a2a_backend="mori")) + self.assertIn("MoE all-to-all", str(ctx.exception)) + + def test_reject_pipeline_parallelism(self): + with self.assertRaises(ValueError) as ctx: + self._run(self._sa(pp_size=2)) + self.assertIn("pipeline parallelism", str(ctx.exception)) + + def test_reject_data_parallelism(self): + with self.assertRaises(ValueError) as ctx: + self._run(self._sa(dp_size=2)) + self.assertIn("data parallelism", str(ctx.exception)) + + def test_reject_decode_context_parallelism(self): + with self.assertRaises(ValueError) as ctx: + self._run(self._sa(dcp_size=2)) + self.assertIn("decode context parallelism", str(ctx.exception)) + + def test_reject_speculative_decoding(self): + with self.assertRaises(ValueError) as ctx: + self._run(self._sa(speculative_algorithm="EAGLE")) + self.assertIn("speculative decoding", str(ctx.exception)) + + def test_no_role_switch_is_unaffected(self): + # The same unsupported feature is fine when role switch is off. + self._run(self._sa(enable_pd_role_switch=False, moe_a2a_backend="mori")) + + +# --- teardown: transfer-worker thread-leak fix + prefix-cache release (radix ON) --- +import threading # noqa: E402 +import time # noqa: E402 + +import zmq # noqa: E402 + +try: + from sglang.srt.disaggregation.common.utils import FastQueue # noqa: E402 + from sglang.srt.disaggregation.mori.conn import MoriKVManager # noqa: E402 + + _HAS_MORI = True +except Exception: # pragma: no cover - environment dependent + _HAS_MORI = False + +try: + from sglang.srt.disaggregation.common.utils import ( # noqa: E402,F811 + FastQueue as _FQ, + ) + from sglang.srt.disaggregation.mooncake.conn import MooncakeKVManager # noqa: E402 + + _HAS_MOONCAKE = True +except Exception: # pragma: no cover - environment dependent + _HAS_MOONCAKE = False + +try: + from sglang.srt.disaggregation.role_switch import ( # noqa: E402 + _release_prefix_cache_for_role_switch, + teardown_disaggregation, + ) + + _HAS_ROLE_SWITCH = True +except Exception: # pragma: no cover - environment dependent + _HAS_ROLE_SWITCH = False + + +@unittest.skipUnless(_HAS_MORI, "mori not importable in this environment") +class TestMoriTeardownNoThreadLeak(unittest.TestCase): + """teardown() must stop+join the transfer workers it started, so a P->D->P + flip loop does not leak _num_shards transfer threads per cycle.""" + + def test_teardown_joins_transfer_workers(self): + m = MoriKVManager.__new__(MoriKVManager) + m.disaggregation_mode = DisaggregationMode.PREFILL + m._stopped = False + m._worker_threads = [] + m._transfer_queues = [FastQueue() for _ in range(3)] + m.server_socket = MagicMock() + m._zmq_ctx = MagicMock() + m.engine = MagicMock() + m.kv_mem_descs = m.aux_mem_descs = m.state_mem_descs = [] + for q in m._transfer_queues: + t = threading.Thread(target=m._transfer_worker, args=(q,), daemon=True) + t.start() + m._worker_threads.append(t) + started = list(m._worker_threads) + time.sleep(0.05) # let workers park in FastQueue.get() + for t in started: + self.assertTrue(t.is_alive()) + + MoriKVManager.teardown(m) + + for t in started: + self.assertFalse(t.is_alive(), "transfer worker survived teardown (leak)") + self.assertEqual(m._worker_threads, []) + self.assertEqual(m._transfer_queues, []) + + +@unittest.skipUnless(_HAS_MOONCAKE, "mooncake not importable in this environment") +class TestMooncakeTeardownNoThreadLeak(unittest.TestCase): + """teardown() must stop+join the transfer workers it started, so a P->D->P + flip loop does not leak transfer threads per cycle.""" + + def test_teardown_joins_transfer_workers(self): + m = MooncakeKVManager.__new__(MooncakeKVManager) + m.disaggregation_mode = DisaggregationMode.PREFILL + m._stopped = False + m.enable_trace = False + m._worker_threads = [] + m.transfer_queues = [_FQ() for _ in range(3)] + m.executors = [concurrent.futures.ThreadPoolExecutor(1) for _ in range(3)] + m.server_socket = MagicMock() + m._zmq_ctx = MagicMock() + m._socket_lock = threading.Lock() + m._socket_cache = {} + m._monitor_cache = {} + m.engine = MagicMock() + # Built from KVArgs' own annotations: teardown walks several ptr/len + # pairs, and hardcoding them here goes stale every time one is added. + from sglang.srt.disaggregation.base.conn import KVArgs + + m.kv_args = SimpleNamespace(**{name: [] for name in KVArgs.__annotations__}) + for i, (q, ex) in enumerate(zip(m.transfer_queues, m.executors)): + t = threading.Thread( + target=m.transfer_worker, args=(q, ex, None, i), daemon=True + ) + t.start() + m._worker_threads.append(t) + started = list(m._worker_threads) + time.sleep(0.05) # let workers park in FastQueue.get() + for t in started: + self.assertTrue(t.is_alive()) + + MooncakeKVManager.teardown(m) + + for t in started: + self.assertFalse(t.is_alive(), "transfer worker survived teardown (leak)") + self.assertEqual(m._worker_threads, []) + self.assertEqual(m.transfer_queues, []) + self.assertEqual(m.executors, []) + + +@unittest.skipUnless(_HAS_MOONCAKE, "mooncake not importable in this environment") +class TestMooncakeBootstrapThreadRobustness(unittest.TestCase): + """The prefill bootstrap loop moved from a blocking recv_multipart() to a + 500ms poll + _stopped check (so teardown, i.e. a runtime role switch, can + stop it). That loop runs on every mooncake PD instance, so pin the + contract with real ZMQ traffic driven through the ABORT -> ABORT_ACK + path: no message loss while idle or bursting, and prompt exit once + _stopped is set. Unlike mori, the loop has no try/except around recv: a + recv error terminates the thread (see test_recv_error_kills_thread). + """ + + class _FlakySocket(zmq.Socket): + """Real PULL socket whose next recv can be forced to fail once, + emulating a transient ZMQ error between poll() and recv().""" + + fail_next_recv = False + + def recv_multipart(self, *args, **kwargs): + if type(self).fail_next_recv: + type(self).fail_next_recv = False + raise RuntimeError("transient recv failure") + return super().recv_multipart(*args, **kwargs) + + def setUp(self): + self._FlakySocket.fail_next_recv = False + self._ctx = zmq.Context() + sock = self._FlakySocket(self._ctx, zmq.PULL) + port = sock.bind_to_random_port("tcp://127.0.0.1") + m = MooncakeKVManager.__new__(MooncakeKVManager) + m._stopped = False + m._worker_threads = [] + m.server_socket = sock + # The receive path is gated on this flag: role switch must be on for + # the poll-with-timeout loop these tests exercise. + m.server_args = SimpleNamespace(enable_pd_role_switch=True) + # Read by the receive loop; off keeps these tests on the plain ACK path. + m.enable_deferred_decode_kv_release = False + # ABORT for an unknown room takes the "ignoring" branch and still + # ACKs, giving a side-effect-free probe of the receive loop. + m.request_status = {} + m._socket_send_locks = {} + + def _connect(endpoint, is_ipv6=False): + m._socket_send_locks.setdefault(endpoint, threading.Lock()) + return m._connect.return_value + + m._connect = MagicMock(side_effect=_connect) + self.m = m + self._push = self._ctx.socket(zmq.PUSH) + self._push.connect(f"tcp://127.0.0.1:{port}") + + def tearDown(self): + self.m._stopped = True + for t in self.m._worker_threads: + t.join(timeout=3.0) + self._push.close(linger=0) + self.m.server_socket.close(linger=0) + self._ctx.destroy(linger=0) + + def _start(self): + MooncakeKVManager.start_prefill_thread(self.m) + (thread,) = self.m._worker_threads + return thread + + def _send_abort(self, room): + self._push.send_multipart( + [b"ABORT", str(room).encode("ascii"), b"127.0.0.1", b"9999"] + ) + + def _wait_acks(self, n, timeout=10.0): + send = self.m._connect.return_value.send_multipart + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + if send.call_count >= n: + return + time.sleep(0.02) + self.fail(f"expected {n} ABORT_ACKs, got {send.call_count}") + + def test_messages_processed_across_idle_poll_timeouts(self): + self._start() + self._send_abort(1) + self._wait_acks(1) + # Idle past a full poll timeout, then traffic must still flow: the + # empty-poll -> continue path must not disturb the socket. + time.sleep(0.8) + self._send_abort(2) + self._wait_acks(2) + + def test_no_message_loss_under_burst(self): + self._start() + n = 200 + for i in range(n): + self._send_abort(i) + # Two-step poll+recv must consume every queued message exactly once. + self._wait_acks(n) + + def test_recv_error_kills_thread(self): + # No try/except guards recv() in the mooncake loop (unlike mori): a + # recv error terminates the thread and the loop stops processing. + # Pin that contract so adding error handling stays a deliberate, + # reviewed change rather than a silent behavior shift. + thread = self._start() + self._FlakySocket.fail_next_recv = True + self._send_abort(3) + thread.join(timeout=2.0) + self.assertFalse(thread.is_alive(), "bootstrap thread survived recv error") + self.assertFalse(self._FlakySocket.fail_next_recv) # fault consumed + + def test_exits_promptly_when_stopped_while_idle(self): + thread = self._start() + self.m._stopped = True + # Poll timeout is 500ms, so the flag must be observed within ~1 cycle + # (this is what keeps teardown / role switch from hanging). + thread.join(timeout=2.0) + self.assertFalse(thread.is_alive(), "bootstrap thread leaked past stop") + + +def _radix_scheduler(disable_radix_cache): + s = MagicMock() + s.disable_radix_cache = disable_radix_cache + tree = MagicMock() + del tree.clear_storage_backend # plain RadixCache has none + s.tree_cache = tree + s.req_to_token_pool = MagicMock() + s.token_to_kv_pool_allocator = MagicMock() + return s + + +@unittest.skipUnless(_HAS_ROLE_SWITCH, "role_switch not importable in this env") +class TestReleasePrefixCacheOnRoleSwitch(unittest.TestCase): + """The flip may run with radix cache ENABLED: teardown resets the tree cache + + KV pools when radix is on, and is a no-op on the historical chunk-cache path.""" + + def test_noop_when_radix_disabled(self): + s = _radix_scheduler(disable_radix_cache=True) + _release_prefix_cache_for_role_switch(s) + s.tree_cache.reset.assert_not_called() + s.token_to_kv_pool_allocator.clear.assert_not_called() + + def test_releases_when_radix_enabled(self): + s = _radix_scheduler(disable_radix_cache=False) + _release_prefix_cache_for_role_switch(s) + s.tree_cache.reset.assert_called_once_with() + s.req_to_token_pool.clear.assert_called_once_with() + s.token_to_kv_pool_allocator.clear.assert_called_once_with() + + def test_teardown_invokes_release(self): + s = _radix_scheduler(disable_radix_cache=False) + s.disaggregation_mode = DisaggregationMode.PREFILL + s.disagg_prefill_bootstrap_queue = None # no queue -> skip km.teardown() + teardown_disaggregation(s) + self.assertIsNone(s.disagg_metadata_buffers) + s.tree_cache.reset.assert_called_once_with() + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/entrypoints/test_server_info.py b/test/registered/unit/entrypoints/test_server_info.py index 60ee04b80..b40d24a01 100644 --- a/test/registered/unit/entrypoints/test_server_info.py +++ b/test/registered/unit/entrypoints/test_server_info.py @@ -69,6 +69,7 @@ class TestModelInfoSerialization(CustomTestCase): "load_format": _CustomModelLoader, "reasoning_parser": None, "tool_call_parser": None, + "disaggregation_mode": "null", } tokenizer_manager = SimpleNamespace( model_config=SimpleNamespace( @@ -97,6 +98,7 @@ class TestModelInfoSerialization(CustomTestCase): reset_context() self.assertEqual(payload["load_format"], f"{__name__}._CustomModelLoader") + self.assertEqual(payload["disaggregation_mode"], "null") json.dumps(payload)