From 2a46cf2ca0fe284a4ca7cc5b0c35cf0a6ecafe43 Mon Sep 17 00:00:00 2001 From: jambow0320 <1193785411@qq.com> Date: Fri, 11 Sep 2026 23:16:17 +0800 Subject: [PATCH] [PD] Share the prefill->decode failure notification across backends (#36612) Co-authored-by: inkcherry Co-authored-by: Shangming Cai --- .../sglang/srt/disaggregation/common/conn.py | 255 ++++++++++++++++-- .../disaggregation/common/staging_handler.py | 8 +- .../srt/disaggregation/mooncake/conn.py | 142 +++------- python/sglang/srt/disaggregation/mori/conn.py | 211 ++------------- python/sglang/srt/disaggregation/nixl/conn.py | 189 +++++++++---- .../disaggregation/test_nixl_backend_basic.py | 16 +- 6 files changed, 453 insertions(+), 368 deletions(-) diff --git a/python/sglang/srt/disaggregation/common/conn.py b/python/sglang/srt/disaggregation/common/conn.py index 5ebcaf811..b2864ae4b 100644 --- a/python/sglang/srt/disaggregation/common/conn.py +++ b/python/sglang/srt/disaggregation/common/conn.py @@ -144,6 +144,19 @@ class PrefillRankInfo: class CommonKVManager(BaseKVManager): + # Wire layout of the prefill->decode terminal status message. The legacy + # layout (mooncake, and ascend which inherits it) is three untagged frames + # ``[room, status, prefill_rank]``; backends whose control socket also + # carries tagged messages prefix a tag frame and may append a reason: + # ``[tag, room, status, prefill_rank, reason]``. + kv_status_msg_tag: Optional[bytes] = None + kv_status_msg_carries_reason: bool = False + + # Used by decode when the prefill reported Failed without a reason frame. + DEFAULT_PREFILL_FAILURE_REASON = ( + "Failed to get kvcache from prefill instance, it might be dead" + ) + def __init__( self, args: KVArgs, @@ -223,7 +236,6 @@ class CommonKVManager(BaseKVManager): self._socket_lock = threading.Lock() self.failure_records: Dict[int, str] = {} self.failure_lock = threading.Lock() - if self.disaggregation_mode == DisaggregationMode.PREFILL: # When SGLANG_DISAGGREGATION_ALL_CP_RANKS_TRANSFER is True, all CP ranks # participate in KV transfer; Otherwise only CP rank 0 sends. @@ -253,6 +265,7 @@ class CommonKVManager(BaseKVManager): self.bootstrap_timeout = envs.SGLANG_DISAGGREGATION_BOOTSTRAP_TIMEOUT.get() elif self.disaggregation_mode == DisaggregationMode.DECODE: self.enable_staging: bool = False + self._staging_handler = None self.connection_pool: Dict[str, Dict[str, Union[str, int]]] = {} self.connection_lock = threading.Lock() self.required_prefill_response_num_table: Dict[int, int] = {} @@ -363,26 +376,233 @@ class CommonKVManager(BaseKVManager): return self.request_status[bootstrap_room] def update_status(self, bootstrap_room: int, status: KVPoll): - if bootstrap_room not in self.request_status: - # Do not resurrect a cleared entry with Failed: once clear() has - # popped the room from request_status, any late update_status(Failed) - # (e.g. from abort()) must be a no-op. Otherwise a Failed entry could - # pollute a future request that reuses the same bootstrap_room. - if status == KVPoll.Failed: - return - self.request_status[bootstrap_room] = status - else: - if status == KVPoll.Failed: - self.request_status[bootstrap_room] = KVPoll.Failed - else: - self.request_status[bootstrap_room] = max( - self.request_status[bootstrap_room], status - ) + current = self.request_status.get(bootstrap_room) + if current is None: + # The room does not exist yet, or clear() already popped it. Only a + # request's opening status may create it: Bootstrapping normally, or + # WaitingForInput for a dummy CP rank (see CommonKVSender.__init__). + # Anything else would resurrect a concluded room and pollute a later + # request that reuses the same bootstrap_room. + if status in (KVPoll.Bootstrapping, KVPoll.WaitingForInput): + self.request_status[bootstrap_room] = status + return + if status == KVPoll.Failed: + self.request_status[bootstrap_room] = KVPoll.Failed + return + if current == KVPoll.Failed: + # Failed is terminal. It also sorts lowest, so the max() below would + # happily promote it back to Transferring or Success. + return + self.request_status[bootstrap_room] = max(current, status) def record_failure(self, bootstrap_room: int, failure_reason: str): with self.failure_lock: self.failure_records[bootstrap_room] = failure_reason + def _room_notify_targets(self, bootstrap_room: int) -> List[Tuple[str, int]]: + infos = self.transfer_infos.get(bootstrap_room) + if not infos: + return [] + # Every non-dummy endpoint, not just the one a caller failed on: the + # others never receive the room's remaining chunks either. + targets: List[Tuple[str, int]] = [] + # Snapshot: the control thread can register a late peer for this room + # while we walk it, and iterating the live view would then raise. + for info in list(infos.values()): + if info.is_dummy: + continue + target = (info.endpoint, info.dst_port) + if target not in targets: + targets.append(target) + return targets + + def _encode_kv_status_message( + self, + *, + bootstrap_room: int, + status: KVPoll, + failure_reason: Optional[str], + ) -> List[bytes]: + parts = [ + str(bootstrap_room).encode("ascii"), + str(int(status)).encode("ascii"), + str(self._prefill_unique_rank()).encode("ascii"), + ] + if self.kv_status_msg_carries_reason: + parts.append((failure_reason or "").encode("utf-8")) + if self.kv_status_msg_tag is not None: + parts.insert(0, self.kv_status_msg_tag) + return parts + + def parse_kv_status_message( + self, msg: List[bytes] + ) -> Optional[Tuple[int, int, int, Optional[str]]]: + """Decode a prefill status message, or None when it is not one.""" + if self.kv_status_msg_tag is not None: + if not msg or msg[0] != self.kv_status_msg_tag: + return None + msg = msg[1:] + if len(msg) < 3: + logger.warning( + "Dropping malformed prefill status message with %d frames", len(msg) + ) + return None + try: + bootstrap_room = int(msg[0].decode("ascii")) + status = int(msg[1].decode("ascii")) + prefill_rank = int(msg[2].decode("ascii")) + except (UnicodeDecodeError, ValueError): + logger.warning("Dropping unparsable prefill status message") + return None + failure_reason = ( + msg[3].decode("utf-8", errors="replace") + if len(msg) > 3 and msg[3] + else None + ) + return bootstrap_room, status, prefill_rank, failure_reason + + def send_kv_status_message( + self, + *, + targets: List[Tuple[str, int]], + bootstrap_room: int, + status: KVPoll, + failure_reason: Optional[str] = None, + ) -> None: + """Push of a terminal transfer status to decode endpoints.""" + if not targets: + return + parts = self._encode_kv_status_message( + bootstrap_room=bootstrap_room, + status=status, + failure_reason=failure_reason, + ) + for endpoint, dst_port in targets: + na = NetworkAddress(endpoint, dst_port) + try: + self._send_multipart_locked(na.to_tcp(), parts, is_ipv6=na.is_ipv6) + except Exception as e: + logger.warning( + f"Failed to sync status {status} of room {bootstrap_room} to " + f"{na.to_host_port_str()}: {e}" + ) + + def conclude_transfer( + self, + *, + bootstrap_room: int, + status: KVPoll, + targets: Optional[List[Tuple[str, int]]] = None, + failure_reason: Optional[str] = None, + ) -> Optional[KVPoll]: + """Returns the status emitted, or None for a cleared room. + + Runs more than once for a room when a staging chunk is deferred past the + last one. ``targets`` defaults to the room's non-dummy decode endpoints. + """ + if bootstrap_room not in self.request_status: + # The sender already cleared this room. Concluding now would + # re-create it in request_status and leave a failure record that a + # request reusing this bootstrap_room would adopt as its own. + return None + if status == KVPoll.Success: + with self.failure_lock: + recorded = self.failure_records.get(bootstrap_room) + if recorded is not None: + status = KVPoll.Failed + failure_reason = recorded + elif self.request_status.get(bootstrap_room) == KVPoll.Failed: + status = KVPoll.Failed + failure_reason = ( + failure_reason or "Room marked Failed before the transfer ended" + ) + if status == KVPoll.Failed: + with self.failure_lock: + # Keep the first root cause; later callers see the symptom. + failure_reason = self.failure_records.setdefault( + bootstrap_room, failure_reason or "KV transfer failed" + ) + + if targets is None: + targets = self._room_notify_targets(bootstrap_room) + self.update_status(bootstrap_room, status) + self.send_kv_status_message( + targets=targets, + bootstrap_room=bootstrap_room, + status=status, + failure_reason=failure_reason, + ) + return status + + def conclude_failure( + self, + *, + bootstrap_room: int, + failure_reason: str, + targets: Optional[List[Tuple[str, int]]] = None, + ) -> Optional[KVPoll]: + """Record the reason, mark the room Failed and tell decode.""" + return self.conclude_transfer( + bootstrap_room=bootstrap_room, + status=KVPoll.Failed, + targets=targets, + failure_reason=failure_reason, + ) + + def apply_prefill_status( + self, + *, + bootstrap_room: int, + status: int, + prefill_rank: int, + failure_reason: Optional[str] = None, + ) -> None: + """Decode-side handling of one prefill rank's terminal status.""" + if bootstrap_room not in self.request_status: + # The room concluded and was cleared. Recording a failure now would + # leave an entry that a later request reusing this bootstrap_room + # would pick up as its own root cause. + logger.debug("Dropping late status for cleared room %s", bootstrap_room) + return + if status == KVPoll.Success: + self.prefill_response_tracker[bootstrap_room].add(prefill_rank) + expected_response_num = self.required_prefill_response_num_table.get( + bootstrap_room + ) + if expected_response_num is None: + logger.warning( + "No expected prefill response count for room %s, prefill rank %s", + bootstrap_room, + prefill_rank, + ) + return + if ( + len(self.prefill_response_tracker[bootstrap_room]) + < expected_response_num + ): + return + # Tell the staging handler no more chunks are coming, before any + # poller can see Success. Only mooncake gets here: NIXL arms the + # handler from its own notifications, mori has no staging. + if self.enable_staging and self._staging_handler is not None: + handler = self._staging_handler + if handler.is_staging_room(bootstrap_room): + handler.submit_last_scatter_async(bootstrap_room) + self.update_status(bootstrap_room, KVPoll.Success) + return + if status == KVPoll.Failed: + self.record_failure( + bootstrap_room, failure_reason or self.DEFAULT_PREFILL_FAILURE_REASON + ) + self.update_status(bootstrap_room, KVPoll.Failed) + return + logger.warning( + "Ignoring non-terminal status %s for room %s from prefill rank %s", + status, + bootstrap_room, + prefill_rank, + ) + def register_deferred_abort_room(self, bootstrap_room: int) -> None: """Arm drain-ack accounting for a held room; a fresh set wipes stale acks from a prior request that reused this bootstrap_room.""" @@ -1613,6 +1833,9 @@ class CommonKVReceiver(BaseKVReceiver): self.kv_mgr.request_status.pop(self.bootstrap_room, None) self.kv_mgr.required_prefill_response_num_table.pop(self.bootstrap_room, None) self.kv_mgr.prefill_response_tracker.pop(self.bootstrap_room, None) + self.kv_mgr.addr_to_rooms_tracker[self.bootstrap_addr].discard( + self.bootstrap_room + ) def abort(self): self.kv_mgr.record_failure( diff --git a/python/sglang/srt/disaggregation/common/staging_handler.py b/python/sglang/srt/disaggregation/common/staging_handler.py index a9ae31b85..ff2f0b6d5 100644 --- a/python/sglang/srt/disaggregation/common/staging_handler.py +++ b/python/sglang/srt/disaggregation/common/staging_handler.py @@ -951,13 +951,7 @@ def prefetch_staging_reqs( full_chunk_pages = staging_grid_tokens(chunked_prefill_size, page_size) // page_size for session_id, tinfo in transfer_infos[room].items(): - # mooncake exposes is_dummy as a dataclass bool field, NIXL exposes it - # as a method (it consults decode_prefix_len). Normalize via callable() - # so this shared helper works for either backend; treating a bound - # method as truthy (the previous behavior) silently dropped every - # STAGING_REQ on NIXL and deadlocked the prefill transfer worker. - is_dummy_attr = tinfo.is_dummy - if is_dummy_attr() if callable(is_dummy_attr) else is_dummy_attr: + if tinfo.is_dummy: continue total_pages = len(tinfo.dst_kv_indices) if total_pages == 0: diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py index a37600e44..970824f99 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -292,7 +292,6 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager): self._staging_ctx = DecodeStagingContext() if self.enable_staging else None if self.enable_staging: self._init_staging_allocator() - self._staging_handler = None self.start_decode_thread() def init_engine(self): @@ -1742,20 +1741,6 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager): return self._transfer_data(req.mooncake_session_id, transfer_blocks) - def sync_status_to_decode_endpoint( - self, remote: str, dst_port: int, room: int, status: int, prefill_rank: int - ): - na = NetworkAddress(remote, dst_port) - self._send_multipart_locked( - na.to_tcp(), - [ - str(room).encode("ascii"), - str(status).encode("ascii"), - str(prefill_rank).encode("ascii"), - ], - is_ipv6=na.is_ipv6, - ) - def transfer_worker( self, queue: FastQueue, @@ -1822,11 +1807,7 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager): polls = [] dst_ranks_infos = [] # Unique id per prefill sender so decode's response set size matches expected_response_num. - prefill_unique_rank = ( - self.attn_tp_rank * (self.pp_size * self.attn_cp_size) - + self.pp_rank * self.attn_cp_size - + self.attn_cp_rank - ) + prefill_unique_rank = self._prefill_unique_rank() # When staging transfer is not yet ready (watermark/allocation pending), # the chunk is re-enqueued and we break out of the req loop to retry later. staging_deferred = False @@ -1836,17 +1817,13 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager): # Early exit if the request has failed with self.session_lock: if req.mooncake_session_id in self.failed_sessions: - self.record_failure( - kv_chunk.room, - f"Decode instance could be dead, remote mooncake session {req.mooncake_session_id} is not alive", - ) - self.update_status(kv_chunk.room, KVPoll.Failed) - self.sync_status_to_decode_endpoint( - req.endpoint, - req.dst_port, - req.room, - KVPoll.Failed, - prefill_unique_rank, + self.conclude_failure( + bootstrap_room=kv_chunk.room, + failure_reason=( + "Decode instance could be dead, remote " + f"mooncake session {req.mooncake_session_id} " + "is not alive" + ), ) break @@ -1989,18 +1966,12 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager): logger.error( f"Session {req.mooncake_session_id} failed." ) - self.record_failure( - kv_chunk.room, - f"Failed to send kv chunk of {kv_chunk.room} to " - f"{NetworkAddress(req.endpoint, req.dst_port).to_host_port_str()}", - ) - self.update_status(kv_chunk.room, KVPoll.Failed) - self.sync_status_to_decode_endpoint( - req.endpoint, - req.dst_port, - req.room, - KVPoll.Failed, - prefill_unique_rank, + self.conclude_failure( + bootstrap_room=kv_chunk.room, + failure_reason=( + f"Failed to send kv chunk of {kv_chunk.room} to " + f"{NetworkAddress(req.endpoint, req.dst_port).to_host_port_str()}" + ), ) break @@ -2020,18 +1991,13 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager): self.failed_sessions.add( req.mooncake_session_id ) - self.record_failure( - kv_chunk.room, - f"Failed to send state components of {kv_chunk.room} to " - f"{NetworkAddress(req.endpoint, req.dst_port).to_host_port_str()}", - ) - self.update_status(kv_chunk.room, KVPoll.Failed) - self.sync_status_to_decode_endpoint( - req.endpoint, - req.dst_port, - req.room, - KVPoll.Failed, - prefill_unique_rank, + self.conclude_failure( + bootstrap_room=kv_chunk.room, + failure_reason=( + "Failed to send state components of " + f"{kv_chunk.room} to " + f"{NetworkAddress(req.endpoint, req.dst_port).to_host_port_str()}" + ), ) break @@ -2042,22 +2008,21 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager): target_rank_registration_info.dst_aux_ptrs, ) polls.append(True if ret == 0 else False) - dst_ranks_infos.append( - (req.endpoint, req.dst_port, req.room) - ) + dst_ranks_infos.append((req.endpoint, req.dst_port)) # Only sync status when all the dst ranks have received the kvcache if len(polls) == req.required_dst_info_num: status = KVPoll.Success if all(polls) else KVPoll.Failed - self.update_status(req.room, status) - for endpoint, dst_port, room in dst_ranks_infos: - self.sync_status_to_decode_endpoint( - endpoint, - dst_port, - room, - status, - prefill_unique_rank, - ) + self.conclude_transfer( + bootstrap_room=req.room, + status=status, + targets=dst_ranks_infos, + failure_reason=( + None + if status == KVPoll.Success + else f"Failed to send aux data of {req.room}" + ), + ) else: # Dummy request means the decode instance is not used, so its status can be marked as success directly # Dummy request does not need to sync status to decode endpoint @@ -2295,32 +2260,16 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager): ) continue - bootstrap_room, status, prefill_rank = msg - status = int(status.decode("ascii")) - bootstrap_room = int(bootstrap_room.decode("ascii")) - prefill_rank = int(prefill_rank.decode("ascii")) - - if status == KVPoll.Success: - if bootstrap_room in self.request_status: - self.prefill_response_tracker[bootstrap_room].add(prefill_rank) - expected_response_num = ( - self.required_prefill_response_num_table[bootstrap_room] - ) - arrived_response_num = len( - self.prefill_response_tracker[bootstrap_room] - ) - if arrived_response_num == expected_response_num: - if self.enable_staging: - handler = self._staging_handler - if handler.is_staging_room(bootstrap_room): - handler.submit_last_scatter_async(bootstrap_room) - self.update_status(bootstrap_room, KVPoll.Success) - elif status == KVPoll.Failed: - self.record_failure( - bootstrap_room, - "Failed to get kvcache from prefill instance, it might be dead", - ) - self.update_status(bootstrap_room, status) + parsed = self.parse_kv_status_message(msg) + if parsed is None: + continue + room, status, prefill_rank, reason = parsed + self.apply_prefill_status( + bootstrap_room=room, + status=status, + prefill_rank=prefill_rank, + failure_reason=reason, + ) threading.Thread(target=decode_thread).start() self._start_heartbeat_checker_thread() @@ -2380,13 +2329,6 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager): def get_session_id(self): return self.engine.get_session_id() - def _on_heartbeat_success(self, bootstrap_addr: str): - current_rooms = self.addr_to_rooms_tracker[bootstrap_addr].copy() - for bootstrap_room in current_rooms: - # Remove KVPoll.Success requests from the tracker - if bootstrap_room not in self.request_status: - self.addr_to_rooms_tracker[bootstrap_addr].discard(bootstrap_room) - def _run_one_probe_pass(self) -> None: with self.session_lock: snapshot = list(self.failed_sessions) diff --git a/python/sglang/srt/disaggregation/mori/conn.py b/python/sglang/srt/disaggregation/mori/conn.py index dee5afb95..42c3409eb 100644 --- a/python/sglang/srt/disaggregation/mori/conn.py +++ b/python/sglang/srt/disaggregation/mori/conn.py @@ -7,7 +7,7 @@ import struct import threading import time import uuid -from typing import Dict, List, Optional, Tuple +from typing import List, Optional import msgspec import numpy as np @@ -294,6 +294,11 @@ class TransferTarget: class MoriKVManager(CommonKVManager): AUX_DATA_HEADER = b"AUX_DATA" + # The bootstrap socket carries several message kinds, so the status message + # is tagged. Mori has always shipped the failure reason with it. + kv_status_msg_tag = MORI_GUARD + kv_status_msg_carries_reason = True + def __init__( self, args: KVArgs, @@ -319,8 +324,6 @@ 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() - self._room_status_notified: Dict[int, bool] = {} - self._room_notify_lock = threading.Lock() for shard, queue in enumerate(self._transfer_queues): threading.Thread( target=self._transfer_worker, @@ -333,7 +336,6 @@ class MoriKVManager(CommonKVManager): ).start() self._start_bootstrap_thread() elif self.disaggregation_mode == DisaggregationMode.DECODE: - self.room_to_bootstrap_addr: Dict[int, str] = {} self._start_decode_thread() self._start_heartbeat_checker_thread() @@ -415,19 +417,6 @@ class MoriKVManager(CommonKVManager): component_descs.append(desc) self.state_mem_descs.append(component_descs) - def update_status(self, bootstrap_room: int, status: KVPoll): - current = self.request_status.get(bootstrap_room) - if current is None: - # Room not yet created or already cleared. - # Only allow initial creation: Bootstrapping (normal) or - # WaitingForInput (dummy CP rank, see CommonKVSender.__init__). - if status not in (KVPoll.Bootstrapping, KVPoll.WaitingForInput): - return - elif current == KVPoll.Failed and status != KVPoll.Failed: - # Failed is terminal — never overwrite with non-Failed. - return - super().update_status(bootstrap_room, status) - def _transfer_worker(self, queue: FastQueue) -> None: while True: kv_chunk = queue.get() @@ -443,7 +432,9 @@ class MoriKVManager(CommonKVManager): except Exception: pass try: - self._conclude_room_failure(kv_chunk.room, failure_reason) + self.conclude_failure( + bootstrap_room=kv_chunk.room, failure_reason=failure_reason + ) except Exception: try: logger.exception( @@ -464,7 +455,7 @@ class MoriKVManager(CommonKVManager): if self._should_skip_transfer(room): return - statuses, target_infos = self._submit_kv_transfer( + statuses = self._submit_kv_transfer( room, kv_chunk.prefill_kv_indices, kv_chunk.index_slice, @@ -480,14 +471,14 @@ class MoriKVManager(CommonKVManager): if self._should_skip_transfer(room): return if failure_reason is not None: - self._conclude_room_failure(room, failure_reason) + self.conclude_failure(bootstrap_room=room, failure_reason=failure_reason) return if kv_chunk.is_last_chunk: - self._notify_decode_for_room( - room, KVPoll.Success, target_infos=target_infos - ) - self.update_status(room, KVPoll.Success) + # conclude_transfer downgrades to Failed when a failure was recorded + # while this chunk was in flight, and applies the same status locally + # and on the wire. + self.conclude_transfer(bootstrap_room=room, status=KVPoll.Success) def _should_skip_transfer(self, room: int) -> bool: if room not in self.request_status or self.check_status(room) == KVPoll.Failed: @@ -523,60 +514,6 @@ class MoriKVManager(CommonKVManager): return f"KV transfer failed: {status.Message()}" return "KV transfer failed due to unknown reason" - def _notify_decode_for_room( - self, - room: int, - status: KVPoll, - failure_reason: Optional[str] = None, - target_infos: Optional[List[TransferInfo]] = None, - ) -> None: - with self._room_notify_lock: - if room not in self.request_status or self._room_status_notified.get(room): - return - - emitted_status = status - emitted_reason = failure_reason - - if emitted_status == KVPoll.Success: - with self.failure_lock: - recorded = self.failure_records.get(room) - if recorded is not None: - emitted_status = KVPoll.Failed - emitted_reason = recorded - elif self.request_status.get(room) == KVPoll.Failed: - emitted_status = KVPoll.Failed - emitted_reason = ( - emitted_reason or "request marked Failed before notify" - ) - - if emitted_status == KVPoll.Failed: - with self.failure_lock: - self.failure_records.setdefault( - room, emitted_reason or "KV transfer failed" - ) - self.update_status(room, KVPoll.Failed) - - infos = target_infos - if infos is None: - with self.transfer_lock: - room_infos = self.transfer_infos.get(room) - infos = ( - list(room_infos.values()) if room_infos is not None else None - ) - - self._room_status_notified[room] = True - - if infos: - self.notify_decode_status(infos, room, emitted_status, emitted_reason) - - def _conclude_room_failure( - self, room: int, failure_reason: Optional[str] = None - ) -> None: - if failure_reason is None: - with self.failure_lock: - failure_reason = self.failure_records.get(room, "KV transfer failed") - self._notify_decode_for_room(room, KVPoll.Failed, failure_reason) - def add_transfer_request( self, bootstrap_room: int, @@ -772,15 +709,6 @@ class MoriKVManager(CommonKVManager): threading.Thread(target=bootstrap_worker, daemon=True).start() - def _cleanup_room_tracking(self, bootstrap_room: int) -> None: - bootstrap_addr = self.room_to_bootstrap_addr.pop(bootstrap_room, None) - if bootstrap_addr is not None: - rooms = self.addr_to_rooms_tracker.get(bootstrap_addr) - if rooms is not None: - rooms.discard(bootstrap_room) - if not rooms: - self.addr_to_rooms_tracker.pop(bootstrap_addr, None) - def _start_decode_thread(self) -> None: def decode_worker(): while True: @@ -790,97 +718,24 @@ class MoriKVManager(CommonKVManager): self._handle_aux_data(msg) continue - if not msg or msg[0] != MORI_GUARD: + parsed = self.parse_kv_status_message(msg) + if parsed is None: logger.warning( "Received malformed status message on decode worker" ) continue - payload = msg[1:] - if len(payload) < 3: - logger.warning("Incomplete status payload received") - continue - bootstrap_room = int(payload[0].decode("ascii")) - if bootstrap_room not in self.request_status: - logger.debug( - "Dropping late status for cleared room %s", - bootstrap_room, - ) - continue - status_code = int(payload[1].decode("ascii")) - prefill_rank = int(payload[2].decode("ascii")) - failure_reason = ( - payload[3].decode("utf-8") - if len(payload) > 3 and payload[3] - else None + room, status, prefill_rank, reason = parsed + self.apply_prefill_status( + bootstrap_room=room, + status=status, + prefill_rank=prefill_rank, + failure_reason=reason, ) - - if status_code == KVPoll.Success: - tracker = self.prefill_response_tracker[bootstrap_room] - tracker.add(prefill_rank) - expected = self.required_prefill_response_num_table.get( - bootstrap_room, 1 - ) - if len(tracker) >= expected: - self.prefill_response_tracker.pop(bootstrap_room, None) - self.update_status(bootstrap_room, KVPoll.Success) - self._cleanup_room_tracking(bootstrap_room) - elif status_code == KVPoll.Failed: - if failure_reason: - self.record_failure(bootstrap_room, failure_reason) - self.prefill_response_tracker.pop(bootstrap_room, None) - self.update_status(bootstrap_room, KVPoll.Failed) - self._cleanup_room_tracking(bootstrap_room) - else: - logger.warning( - "Unknown status code %s received for room %s", - status_code, - bootstrap_room, - ) except Exception: logger.exception("Decode status worker failed") threading.Thread(target=decode_worker, daemon=True).start() - def _compute_prefill_unique_rank(self) -> int: - """Unique id per prefill sender, encoding TP/PP/CP ranks. - Must match Mooncake's formula so decode's response set size matches - expected_response_num when multiple CP ranks participate.""" - return ( - self.attn_tp_rank * (self.pp_size * self.attn_cp_size) - + self.pp_rank * self.attn_cp_size - + self.attn_cp_rank - ) - - def notify_decode_status( - self, - infos: List[TransferInfo], - bootstrap_room: int, - status: KVPoll, - failure_reason: Optional[str] = None, - ) -> None: - if not infos: - return - payload = [ - MORI_GUARD, - str(bootstrap_room).encode("ascii"), - str(int(status)).encode("ascii"), - str(self._compute_prefill_unique_rank()).encode("ascii"), - failure_reason.encode("utf-8") if failure_reason else b"", - ] - for info in infos: - try: - na = NetworkAddress(info.endpoint, info.dst_port) - socket = self._connect_threadsafe(na.to_tcp(), is_ipv6=na.is_ipv6) - socket.send_multipart(payload) - except Exception: - logger.exception( - "Failed to sync status %s to decode endpoint %s:%s for room %s", - status, - info.endpoint, - info.dst_port, - bootstrap_room, - ) - def _add_remote_peer(self, register_info: KVArgsRegisterInfo) -> None: engine_key = register_info.engine_key if engine_key in self.decode_kv_args_table: @@ -1541,21 +1396,20 @@ class MoriKVManager(CommonKVManager): is_last_chunk: bool, aux_index: Optional[int] = None, state_indices: Optional[List[npt.NDArray[np.int32]]] = None, - ) -> Tuple[List[TransferStatus], Optional[List[TransferInfo]]]: + ) -> List[TransferStatus]: assert self.disaggregation_mode == DisaggregationMode.PREFILL if ( bootstrap_room not in self.request_status or self.request_status.get(bootstrap_room) == KVPoll.Failed ): - return [], None + return [] targets: List[TransferTarget] = [] - target_infos_snapshot: Optional[List[TransferInfo]] = None with self.transfer_lock: current = self.request_status.get(bootstrap_room) if current is None or current == KVPoll.Failed: - return [], None + return [] transfer_infos = self.transfer_infos.get(bootstrap_room) if not transfer_infos: @@ -1571,8 +1425,6 @@ class MoriKVManager(CommonKVManager): f"Peer info missing for engine {info.engine_key}" ) targets.append(TransferTarget(info=info, peer_info=peer_info)) - if is_last_chunk: - target_infos_snapshot = list(transfer_infos.values()) result_statuses: List[TransferStatus] = [] try: @@ -1616,7 +1468,7 @@ class MoriKVManager(CommonKVManager): ) raise RuntimeError(f"Transfer submission failed: {e}") from e - return result_statuses, target_infos_snapshot + return result_statuses class MoriKVSender(CommonKVSender): @@ -1705,11 +1557,6 @@ class MoriKVSender(CommonKVSender): self.conclude_state = status return status - def clear(self) -> None: - super().clear() - with self.kv_mgr._room_notify_lock: - self.kv_mgr._room_status_notified.pop(self.bootstrap_room, None) - def failure_exception(self): if self.conclude_state is None: self.conclude_state = KVPoll.Failed @@ -1741,9 +1588,6 @@ class MoriKVReceiver(CommonKVReceiver): prefill_dp_rank: int, ): super().init(prefill_dp_rank) - if self.bootstrap_room is None: - return - self.kv_mgr.room_to_bootstrap_addr[self.bootstrap_room] = self.bootstrap_addr def _register_kv_args(self) -> bool: if self.bootstrap_infos is None: @@ -1871,7 +1715,6 @@ class MoriKVReceiver(CommonKVReceiver): if self.bootstrap_room is None: return super().clear() - self.kv_mgr._cleanup_room_tracking(self.bootstrap_room) def failure_exception(self): if self.conclude_state is None: diff --git a/python/sglang/srt/disaggregation/nixl/conn.py b/python/sglang/srt/disaggregation/nixl/conn.py index 6bf84a189..980568e2e 100644 --- a/python/sglang/srt/disaggregation/nixl/conn.py +++ b/python/sglang/srt/disaggregation/nixl/conn.py @@ -72,6 +72,12 @@ logger = logging.getLogger(__name__) GUARD = "NixlMsgGuard".encode("ascii") KV_MEM_KINDS = {"VRAM", "DRAM"} +# Once one handle of a batch reports ERR, its siblings settle only when NIXL +# notices their peer is gone, which for UCX means waiting out peer keepalive. +# This worker serves other rooms, so bound that wait. +NIXL_ERR_SETTLE_TIMEOUT_S = 5.0 +NIXL_ERR_SETTLE_POLL_S = 0.001 + def _normalize_kv_mem_kinds(kinds: Optional[List[str]], expected_len: int) -> List[str]: if kinds is None: @@ -166,45 +172,47 @@ class TransferInfo: required_dst_info_num: int dst_state_indices: List[List[int]] decode_prefix_len: Optional[int] = None # for decode radix cache - is_dummy_rank: Optional[bool] = None + is_dummy: bool = False # NOTE: optional staging field; populated via STAGING_RSP. Keep at the # end so positional construction in from_zmq() continues to work. staging: Optional[StagingTransferInfo] = None - def is_dummy(self): - # A transfer is "dummy" only for CP non-authoritative ranks. - # When dst_kv_indices is empty due to a decode-side radix cache - # full hit (decode_prefix_len > 0), the transfer is NOT dummy -- - # aux/state data still needs to be sent. - if self.is_dummy_rank is not None: - return self.is_dummy_rank - if self.dst_kv_indices.size == 0 and self.decode_prefix_len: - return False - return self.dst_kv_indices.size == 0 - @classmethod def from_zmq(cls, msg: List[bytes]): dst_state_indices = ( unpack_int_lists(msg[7], "i") if len(msg) > 7 and msg[7] != b"" else [] ) + dst_kv_indices = np.frombuffer(msg[4], dtype=np.int32) + decode_prefix_len = ( + int(msg[8].decode("ascii")) if len(msg) > 8 and msg[8] != b"" else None + ) # hacky just add it into the message that will be sent + dummy_rank = ( + bool(int(msg[9].decode("ascii"))) + if len(msg) > 9 and msg[9] != b"" + else None + ) + # A transfer is "dummy" only for CP non-authoritative ranks. When + # dst_kv_indices is empty due to a decode-side radix cache full hit + # (decode_prefix_len > 0), the transfer is NOT dummy -- aux/state data + # still needs to be sent. + if dummy_rank is not None: + is_dummy = dummy_rank + elif dst_kv_indices.size == 0 and decode_prefix_len: + is_dummy = False + else: + is_dummy = dst_kv_indices.size == 0 return cls( room=int(msg[0].decode("ascii")), endpoint=msg[1].decode("ascii"), dst_port=int(msg[2].decode("ascii")), agent_name=msg[3].decode("ascii"), - dst_kv_indices=np.frombuffer(msg[4], dtype=np.int32), + dst_kv_indices=dst_kv_indices, dst_aux_index=int(msg[5].decode("ascii")), required_dst_info_num=int(msg[6].decode("ascii")), dst_state_indices=dst_state_indices, - decode_prefix_len=( - int(msg[8].decode("ascii")) if len(msg) > 8 and msg[8] != b"" else None - ), # hacky just add it into the message that will be sent - is_dummy_rank=( - bool(int(msg[9].decode("ascii"))) - if len(msg) > 9 and msg[9] != b"" - else None - ), + decode_prefix_len=decode_prefix_len, + is_dummy=is_dummy, ) @@ -399,6 +407,11 @@ class TransferStatus: class NixlKVManager(StagingManagerMixin, CommonKVManager): + # The decode control socket multiplexes tagged messages, so the status + # message is tagged too. It is new to NIXL, hence free to carry the reason. + kv_status_msg_tag = b"KV_STATUS" + kv_status_msg_carries_reason = True + def __init__( self, args: KVArgs, @@ -523,9 +536,7 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager): ) if self.enable_staging: self._init_staging_decode_ctx() - self._staging_handler = None - if self.enable_staging or self.enable_deferred_decode_kv_release: - self._start_decode_listener_thread() + self._start_decode_listener_thread() self._start_heartbeat_checker_thread() else: raise ValueError( @@ -595,15 +606,19 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager): self._staging_ctx.room_receivers[room] = receiver def _start_decode_listener_thread(self): - """Decode-side ZMQ listener for STAGING_REQ and ABORT_ACK. A thread, not - NIXL notifs: the decode agent has no progress thread, so notifs only drain - inside a live receiver's poll() and would be missed while idle.""" + """Decode-side ZMQ listener for KV_STATUS, STAGING_REQ and ABORT_ACK. A + thread, not NIXL notifs: the decode agent has no progress thread, so notifs + only drain inside a live receiver's poll() and would be missed while idle. + + Started unconditionally: KV_STATUS carries prefill-side transfer failures, + which are independent of staging and deferred KV release.""" def decode_listener_thread(): while True: msg = self.server_socket.recv_multipart() if msg[0] == b"STAGING_REQ": - self._handle_staging_req(msg) + if self.enable_staging: + self._handle_staging_req(msg) continue if msg[0] == b"ABORT_ACK": # Drain ack for an aborted room; aggregate per prefill rank. @@ -612,6 +627,16 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager): int(msg[1].decode("ascii")), int(msg[2].decode("ascii")) ) continue + parsed = self.parse_kv_status_message(msg) + if parsed is not None: + room, status, prefill_rank, reason = parsed + self.apply_prefill_status( + bootstrap_room=room, + status=status, + prefill_rank=prefill_rank, + failure_reason=reason, + ) + continue logger.warning( "decode_listener_thread: unexpected message tag %s", msg[0][:20], @@ -634,7 +659,7 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager): room_infos = self.transfer_infos.get(room, {}) needs_staging = any( - not tinfo.is_dummy() + not tinfo.is_dummy and tinfo.agent_name in self.decode_kv_args_table and self.decode_kv_args_table[tinfo.agent_name].decode_tp_size != self.attn_tp_size @@ -662,11 +687,44 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager): def check_status(self, bootstrap_room: int): return self.request_status.get(bootstrap_room, KVPoll.WaitingForInput) - def update_status(self, bootstrap_room: int, status: KVPoll): - # Keep Failed sticky until the sender clears the room. - if self.request_status.get(bootstrap_room) == KVPoll.Failed: - return - super().update_status(bootstrap_room, status) + def _await_handles( + self, handles: List[Any], *, failure_seen: bool + ) -> Tuple[bool, bool]: + """Poll until every handle settled. Returns ``(settled, any_failed)``. + + The wait is unbounded while every handle is still healthy, and bounded + to NIXL_ERR_SETTLE_TIMEOUT_S from the moment the batch is known broken. + ``failure_seen`` arms that deadline up front, for a batch that raised + before the barrier ran. A state that cannot be read counts as running, + since it does not prove the write into the decode's pages is over. + """ + deadline = time.time() + NIXL_ERR_SETTLE_TIMEOUT_S if failure_seen else None + while True: + all_settled = True + any_failed = failure_seen + try: + for handle in handles: + state = self.agent.check_xfer_state(handle) + if state == "ERR": + any_failed = True + elif state != "DONE": + all_settled = False + except Exception as e: + logger.warning(f"Failed to read NIXL transfer state: {e}") + return False, True + if all_settled: + return True, any_failed + if not any_failed: + time.sleep(0) + continue + # This room is already lost, so trade its notification for the + # worker's other rooms: back off, and give up waiting for the + # siblings once the deadline passes. + if deadline is None: + deadline = time.time() + NIXL_ERR_SETTLE_TIMEOUT_S + elif time.time() >= deadline: + return False, True + time.sleep(NIXL_ERR_SETTLE_POLL_S) def _prep_equal_tp_dlist( self, @@ -1095,6 +1153,7 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager): kv_chunk: TransferKVChunk = queue.get() room = kv_chunk.room handles: List[Any] = [] + settle_timed_out = False try: if room not in self.request_status: logger.debug( @@ -1154,7 +1213,7 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager): for req in reqs_to_be_processed: assert room == req.room - if req.is_dummy(): + if req.is_dummy: continue assert req.agent_name in self.decode_kv_args_table @@ -1353,19 +1412,19 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager): # Chunk has been re-enqueued; do not advance status. continue - while handles: - all_done = True - for handle in handles: - state = self.agent.check_xfer_state(handle) - if state == "ERR": - raise RuntimeError( - f"NIXL transfer encountered ERR room={room}" - ) - if state != "DONE": - all_done = False - if all_done: - break - time.sleep(0) + # Raise only once every handle of this batch settled, not on the + # first ERR: a sibling still in PROC keeps writing into the + # decode's KV pages, and the failure path below tells the decode + # those pages are free. + settled, any_failed = self._await_handles(handles, failure_seen=False) + if not settled: + settle_timed_out = True + raise RuntimeError( + f"NIXL transfer for room {room} left a handle running " + f"{NIXL_ERR_SETTLE_TIMEOUT_S}s after a peer handle failed" + ) + if any_failed: + raise RuntimeError(f"NIXL transfer encountered ERR room={room}") self._staging_outstanding[room] -= 1 if self.enable_deferred_decode_kv_release: @@ -1409,10 +1468,20 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager): f"Unexpected transfer worker error for room {room}" ) self.exceptions[room] = e - self.record_failure(room, str(e)) - self.update_status(room, KVPoll.Failed) - # No ack here on purpose: the DONE barrier bails on the first - # ERR, so siblings may still be writing; fall back to the timeout. + # An exception raised while the batch was still being built + # leaves the handles posted so far running, so settle here too + # rather than only after the barrier. + notify = False + if not settle_timed_out: + notify, _ = self._await_handles(handles, failure_seen=True) + if notify: + self.conclude_failure(bootstrap_room=room, failure_reason=str(e)) + else: + # A handle can still write into the decode's KV pages, so + # leave the room to the decode's waiting timeout rather + # than telling it those pages are free. + self.record_failure(room, str(e)) + self.update_status(room, KVPoll.Failed) def register_buffer_to_engine(self): self.kv_descs = [] @@ -3002,6 +3071,13 @@ class NixlKVReceiver(CommonKVReceiver): super().__init__(mgr, bootstrap_addr, bootstrap_room) self.init_time = None + def clear(self) -> None: + super().clear() + # transfer_statuses is NIXL's own per-room bookkeeping -- the other + # backends track completion through prefill_response_tracker, which + # CommonKVReceiver.clear() already drops -- so it needs its own pop. + self.kv_mgr.transfer_statuses.pop(self.bootstrap_room, None) + def send_metadata( self, kv_indices: npt.NDArray[np.int32], @@ -3096,11 +3172,7 @@ class NixlKVReceiver(CommonKVReceiver): # deadline would otherwise lose to the timeout purely by poll ordering. self.kv_mgr.update_transfer_status() if self.kv_mgr.check_transfer_done(self.bootstrap_room): # type: ignore - self.kv_mgr.addr_to_rooms_tracker[self.bootstrap_addr].discard( - self.bootstrap_room - ) self.conclude_state = KVPoll.Success - del self.kv_mgr.transfer_statuses[self.bootstrap_room] return self.conclude_state # type: ignore timeout_result = self._check_waiting_timeout() @@ -3201,6 +3273,11 @@ class NixlKVReceiver(CommonKVReceiver): return True def failure_exception(self): + if self.conclude_state is None: + self.conclude_state = KVPoll.Failed + + self.clear() + with self.kv_mgr.failure_lock: failure_reason = self.kv_mgr.failure_records.pop(self.bootstrap_room, None) is_propagated = failure_reason is None diff --git a/test/registered/unit/disaggregation/test_nixl_backend_basic.py b/test/registered/unit/disaggregation/test_nixl_backend_basic.py index 60d5bf3e4..1553b08a9 100644 --- a/test/registered/unit/disaggregation/test_nixl_backend_basic.py +++ b/test/registered/unit/disaggregation/test_nixl_backend_basic.py @@ -165,7 +165,7 @@ class TestNixlTransferInfo(CustomTestCase): ] ) - self.assertFalse(info.is_dummy()) + self.assertFalse(info.is_dummy) def test_empty_indices_without_decode_prefix_is_dummy(self): info = TransferInfo.from_zmq( @@ -182,7 +182,7 @@ class TestNixlTransferInfo(CustomTestCase): ] ) - self.assertTrue(info.is_dummy()) + self.assertTrue(info.is_dummy) def test_explicit_dummy_frame_true_is_dummy(self): # msg[9] is the explicit is_dummy frame the sender writes @@ -848,10 +848,9 @@ class TestNixlReceiverPoll(CustomTestCase): mgr.update_transfer_status.assert_called_once_with() mgr.record_failure.assert_not_called() mgr.update_status.assert_not_called() - self.assertNotIn(11, mgr.transfer_statuses) @patch("sglang.srt.disaggregation.nixl.conn.time.time") - def test_transfer_done_returns_success_and_cleans_room_state(self, mock_time): + def test_transfer_done_returns_success_and_clear_drops_room_state(self, mock_time): mock_time.return_value = 12.0 receiver, mgr = self._make_receiver(status=KVPoll.WaitingForInput) receiver.started_transfer = True @@ -864,9 +863,16 @@ class TestNixlReceiverPoll(CustomTestCase): mgr.check_transfer_done.return_value = True self.assertEqual(receiver.poll(), KVPoll.Success) + self.assertEqual(receiver.conclude_state, KVPoll.Success) + + # poll() only concludes now; dropping room state is left to clear(), the + # way mooncake and mori already do it. The scheduler calls clear() as + # soon as poll() reports Success or Failed, so both terminal paths clean + # up -- the cleanup that used to live in poll() ran on Success only. + receiver.clear() + self.assertNotIn(11, mgr.transfer_statuses) self.assertNotIn(11, mgr.addr_to_rooms_tracker["prefill:8998"]) - self.assertEqual(receiver.conclude_state, KVPoll.Success) class TestNixlNodeFailure(CustomTestCase):