From 96f5066dcf9ba985687d5e00c983548d9b282b97 Mon Sep 17 00:00:00 2001 From: jambow0320 <1193785411@qq.com> Date: Thu, 13 Aug 2026 16:19:53 +1000 Subject: [PATCH 1/9] add nixl bootstrap timeout Co-authored-by: Cursor --- python/sglang/srt/disaggregation/nixl/conn.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/python/sglang/srt/disaggregation/nixl/conn.py b/python/sglang/srt/disaggregation/nixl/conn.py index 97e6251930fa..b8cbe7925d71 100644 --- a/python/sglang/srt/disaggregation/nixl/conn.py +++ b/python/sglang/srt/disaggregation/nixl/conn.py @@ -2745,6 +2745,7 @@ def __init__( pp_rank, req_has_disagg_prefill_dp_rank, ) + self.init_time = time.time() self.has_sent = False self.chunk_id = 0 self._send_failed = False @@ -2790,6 +2791,10 @@ def poll(self) -> KVPoll: if self._send_failed: return KVPoll.Failed # type: ignore status = self.kv_mgr.check_status(self.bootstrap_room) + if status == KVPoll.Bootstrapping: + timeout_result = self._check_bootstrap_timeout() + if timeout_result is not None: + return timeout_result # Hold Success until all staging chunks transferred: a deferred chunk # can still be pending, and concluding now would drop it. if ( From b285ea79ce9f14f1238e3b2b4b0e66d0efbf9f29 Mon Sep 17 00:00:00 2001 From: jambow0320 <1193785411@qq.com> Date: Sun, 16 Aug 2026 18:39:27 +1000 Subject: [PATCH 2/9] align Mori Decode heartbeat --- python/sglang/srt/disaggregation/mori/conn.py | 1 + 1 file changed, 1 insertion(+) diff --git a/python/sglang/srt/disaggregation/mori/conn.py b/python/sglang/srt/disaggregation/mori/conn.py index 7954877f7997..f2b0269ab86c 100644 --- a/python/sglang/srt/disaggregation/mori/conn.py +++ b/python/sglang/srt/disaggregation/mori/conn.py @@ -341,6 +341,7 @@ def __init__( elif self.disaggregation_mode == DisaggregationMode.DECODE: self.room_to_bootstrap_addr: Dict[int, str] = {} self._start_decode_thread() + self._start_heartbeat_checker_thread() def _init_engine(self) -> IOEngine: if self.kv_args.ib_device: From 62c13f66ba85f109ac3eb3bf2dfd63d34cd9d37d Mon Sep 17 00:00:00 2001 From: jambow0320 <1193785411@qq.com> Date: Sun, 16 Aug 2026 19:59:03 +1000 Subject: [PATCH 3/9] align Mori abort handling with common protocol --- python/sglang/srt/disaggregation/mori/conn.py | 33 +++++++++++++++++++ 1 file changed, 33 insertions(+) diff --git a/python/sglang/srt/disaggregation/mori/conn.py b/python/sglang/srt/disaggregation/mori/conn.py index f2b0269ab86c..e913b1cfd577 100644 --- a/python/sglang/srt/disaggregation/mori/conn.py +++ b/python/sglang/srt/disaggregation/mori/conn.py @@ -544,6 +544,37 @@ def _handle_transfer_message(self, payload: List[bytes]) -> None: except Exception: logger.exception("Failed to parse transfer info message") + def _handle_abort_notification(self, msg: List[bytes]) -> bool: + if not msg or msg[0] != b"ABORT": + return False + + try: + room_to_be_aborted = int(msg[1].decode("ascii")) + except Exception as e: + logger.debug(f"Ignoring malformed abort notification: {e}") + return True + + if ( + room_to_be_aborted in self.request_status + and self.check_status(room_to_be_aborted) != KVPoll.Success + ): + self.record_failure( + room_to_be_aborted, + "Aborted by decode-side abort notification.", + ) + self.update_status(room_to_be_aborted, KVPoll.Failed) + logger.debug( + f"Received abort notification for room {room_to_be_aborted}, " + "marked as Failed" + ) + else: + logger.debug( + f"Received abort notification for room {room_to_be_aborted}, " + "ignoring (already completed or unknown)" + ) + + return True + def _validate_message(self, msg: List[bytes]) -> Optional[List[bytes]]: if not msg or msg[0] != MORI_GUARD: logger.warning("Received malformed bootstrap message") @@ -558,6 +589,8 @@ def bootstrap_worker(): while True: try: msg = self.server_socket.recv_multipart() + if self._handle_abort_notification(msg): + continue payload = self._validate_message(msg) if payload is None: continue From 46056556a5f2b9d06321fab43a038bc4a36c72cd Mon Sep 17 00:00:00 2001 From: jambow0320 <1193785411@qq.com> Date: Mon, 17 Aug 2026 13:23:09 +1000 Subject: [PATCH 4/9] skip cleared rooms in NIXL and Mori workers --- python/sglang/srt/disaggregation/mori/conn.py | 2 ++ python/sglang/srt/disaggregation/nixl/conn.py | 10 +++++++++- 2 files changed, 11 insertions(+), 1 deletion(-) diff --git a/python/sglang/srt/disaggregation/mori/conn.py b/python/sglang/srt/disaggregation/mori/conn.py index e913b1cfd577..3d5e92ca90de 100644 --- a/python/sglang/srt/disaggregation/mori/conn.py +++ b/python/sglang/srt/disaggregation/mori/conn.py @@ -1487,6 +1487,8 @@ def _maybe_finalize_if_room_failed(self) -> None: self._finalize_failure() def _run_chunk(self, task: _TransferChunk) -> None: + if self.bootstrap_room not in self.kv_mgr.request_status: + return if self.conclude_state is not None: return if self.kv_mgr.request_status.get(self.bootstrap_room) == KVPoll.Failed: diff --git a/python/sglang/srt/disaggregation/nixl/conn.py b/python/sglang/srt/disaggregation/nixl/conn.py index b8cbe7925d71..04c48ffc80e4 100644 --- a/python/sglang/srt/disaggregation/nixl/conn.py +++ b/python/sglang/srt/disaggregation/nixl/conn.py @@ -1119,7 +1119,15 @@ def transfer_worker(self, queue: FastQueue, staging_buffer=None): room = kv_chunk.room handles: List[Any] = [] try: - if self.check_status(room) == KVPoll.Failed: + if ( + room not in self.request_status + or self.check_status(room) == KVPoll.Failed + ): + logger.debug( + "Skipping chunk for room %s because it has already " + "failed, been aborted, or been cleared", + room, + ) self._staging_outstanding.pop(room, None) continue From a716ceaac7adecb859e1ba1b31ea5e334ddfefc3 Mon Sep 17 00:00:00 2001 From: jambow0320 <1193785411@qq.com> Date: Mon, 17 Aug 2026 20:17:47 +1000 Subject: [PATCH 5/9] harden Mooncake and NIXL control message parsing --- .../srt/disaggregation/mooncake/conn.py | 326 +++++++++--------- python/sglang/srt/disaggregation/nixl/conn.py | 134 +++---- 2 files changed, 241 insertions(+), 219 deletions(-) diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py index 66ac63b34c98..3119785b07e1 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -1961,178 +1961,192 @@ 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() - room = waiting_req_bytes[0].decode("ascii") - # Staging: decode reports consumption watermark back to prefill - if room == "WATERMARK": - from sglang.srt.disaggregation.common.staging_handler import ( - handle_watermark_msg, - ) - - handle_watermark_msg(self._staging_ctx, waiting_req_bytes) - continue - # Staging: decode replies with allocated staging offset - if room == "STAGING_RSP": - from sglang.srt.disaggregation.common.staging_handler import ( - handle_staging_rsp, - ) - - handle_staging_rsp(waiting_req_bytes, self.transfer_infos) - continue - # Decode-side abort notification: mark room as failed and ACK - if room == "ABORT": - room_to_be_aborted = int(waiting_req_bytes[1].decode("ascii")) - decode_ip = waiting_req_bytes[2].decode("ascii") - decode_port = int(waiting_req_bytes[3].decode("ascii")) - # No need to abort the room if it has already succeeded - if ( - room_to_be_aborted in self.request_status - and self.check_status(room_to_be_aborted) != KVPoll.Success - ): - self.update_status(room_to_be_aborted, KVPoll.Failed) - logger.debug( - f"Received abort notification for room {room_to_be_aborted}, " - f"marked as Failed" - ) - else: - logger.debug( - f"Received abort notification for room {room_to_be_aborted}, " - f"ignoring (already completed or unknown)" + try: + waiting_req_bytes = self.server_socket.recv_multipart() + room = waiting_req_bytes[0].decode("ascii") + # Staging: decode reports consumption watermark back to prefill + if room == "WATERMARK": + from sglang.srt.disaggregation.common.staging_handler import ( + handle_watermark_msg, ) - # Send ACK back to decode endpoint - try: - na = NetworkAddress(decode_ip, decode_port) - self._send_multipart_locked( - na.to_tcp(), - [ - b"ABORT_ACK", - str(room_to_be_aborted).encode("ascii"), - ], - is_ipv6=na.is_ipv6, + + handle_watermark_msg(self._staging_ctx, waiting_req_bytes) + continue + # Staging: decode replies with allocated staging offset + if room == "STAGING_RSP": + from sglang.srt.disaggregation.common.staging_handler import ( + handle_staging_rsp, ) - logger.debug( - f"Sent ABORT_ACK for room {room_to_be_aborted} to " - f"{decode_ip}:{decode_port}" + + handle_staging_rsp(waiting_req_bytes, self.transfer_infos) + continue + # Decode-side abort notification: mark room as failed and ACK + if room == "ABORT": + room_to_be_aborted = int(waiting_req_bytes[1].decode("ascii")) + decode_ip = waiting_req_bytes[2].decode("ascii") + decode_port = int(waiting_req_bytes[3].decode("ascii")) + # No need to abort the room if it has already succeeded + if ( + room_to_be_aborted in self.request_status + and self.check_status(room_to_be_aborted) != KVPoll.Success + ): + self.update_status(room_to_be_aborted, KVPoll.Failed) + logger.debug( + f"Received abort notification for room {room_to_be_aborted}, " + f"marked as Failed" + ) + else: + logger.debug( + f"Received abort notification for room {room_to_be_aborted}, " + f"ignoring (already completed or unknown)" + ) + # Send ACK back to decode endpoint + try: + na = NetworkAddress(decode_ip, decode_port) + self._send_multipart_locked( + na.to_tcp(), + [ + b"ABORT_ACK", + str(room_to_be_aborted).encode("ascii"), + ], + is_ipv6=na.is_ipv6, + ) + logger.debug( + f"Sent ABORT_ACK for room {room_to_be_aborted} to " + f"{decode_ip}:{decode_port}" + ) + except Exception as e: + logger.debug( + f"Failed to send ABORT_ACK for room {room_to_be_aborted}: {e}" + ) + continue + mooncake_session_id = waiting_req_bytes[3].decode("ascii") + if room == "None": + decode_kv_args = KVArgsRegisterInfo.from_zmq(waiting_req_bytes) + decode_kv_args.requires_dcp_relayout = ( + self.requires_dcp_relayout( + decode_kv_args.dst_dcp_size, + decode_kv_args.dst_dcp_rank, + ) ) - except Exception as e: + if decode_kv_args.requires_dcp_relayout: + decode_kv_args.dcp_token_item_lens = ( + self.prepare_dcp_token_item_lens( + [decode_kv_args.dst_kv_item_len] + * len(self.kv_args.kv_item_lens) + ) + ) + self.decode_kv_args_table[mooncake_session_id] = decode_kv_args + with self.session_lock: + if mooncake_session_id in self.failed_sessions: + self.failed_sessions.remove(mooncake_session_id) + if mooncake_session_id in self.session_failures: + del self.session_failures[mooncake_session_id] logger.debug( - f"Failed to send ABORT_ACK for room {room_to_be_aborted}: {e}" + f"Register KVArgs from {mooncake_session_id} successfully" ) - continue - mooncake_session_id = waiting_req_bytes[3].decode("ascii") - if room == "None": - decode_kv_args = KVArgsRegisterInfo.from_zmq(waiting_req_bytes) - decode_kv_args.requires_dcp_relayout = self.requires_dcp_relayout( - decode_kv_args.dst_dcp_size, - decode_kv_args.dst_dcp_rank, - ) - if decode_kv_args.requires_dcp_relayout: - decode_kv_args.dcp_token_item_lens = ( - self.prepare_dcp_token_item_lens( - [decode_kv_args.dst_kv_item_len] - * len(self.kv_args.kv_item_lens) - ) + continue + else: + required_dst_info_num = int( + waiting_req_bytes[7].decode("ascii") ) - self.decode_kv_args_table[mooncake_session_id] = decode_kv_args - with self.session_lock: - if mooncake_session_id in self.failed_sessions: - self.failed_sessions.remove(mooncake_session_id) - if mooncake_session_id in self.session_failures: - del self.session_failures[mooncake_session_id] - logger.debug( - f"Register KVArgs from {mooncake_session_id} successfully" - ) - continue - else: - required_dst_info_num = int(waiting_req_bytes[7].decode("ascii")) - room = int(room) - if room not in self.transfer_infos: - self.transfer_infos[room] = {} + room = int(room) + if room not in self.transfer_infos: + self.transfer_infos[room] = {} - self.transfer_infos[room][mooncake_session_id] = ( - TransferInfo.from_zmq(waiting_req_bytes) - ) - # NOTE: after bootstrapping we can mark the req as waiting for input - if len(self.transfer_infos[room]) == required_dst_info_num: - self.resolve_kv_replica_factor(self.transfer_infos[room]) - self.req_to_decode_prefix_len[room] = next( - ( - info.decode_prefix_len - for info in self.transfer_infos[room].values() - if info.decode_prefix_len is not None - ), - 0, + self.transfer_infos[room][mooncake_session_id] = ( + TransferInfo.from_zmq(waiting_req_bytes) ) - self.update_status(room, KVPoll.WaitingForInput) + # NOTE: after bootstrapping we can mark the req as waiting for input + if len(self.transfer_infos[room]) == required_dst_info_num: + self.resolve_kv_replica_factor(self.transfer_infos[room]) + self.req_to_decode_prefix_len[room] = next( + ( + info.decode_prefix_len + for info in self.transfer_infos[room].values() + if info.decode_prefix_len is not None + ), + 0, + ) + self.update_status(room, KVPoll.WaitingForInput) + except Exception: + logger.exception("Bootstrap worker failed") threading.Thread(target=bootstrap_thread).start() def start_decode_thread(self): def decode_thread(): while True: - msg = self.server_socket.recv_multipart() - if msg[0] == MooncakeKVManager.AUX_DATA_HEADER: - self._handle_aux_data(msg) - continue - - # Staging: prefill notifies a chunk written to staging buffer - if msg[0] == b"CHUNK_READY": - room = int(msg[1].decode("ascii")) - chunk_idx = int(msg[2].decode("ascii")) - page_start = int(msg[3].decode("ascii")) - num_pages = int(msg[4].decode("ascii")) - session_id = msg[5].decode("ascii") - handler = self._staging_handler - assert ( - handler is not None - ), "CHUNK_READY received before staging handler initialized" - handler.handle_chunk_arrived( - room, - chunk_idx, - page_start, - num_pages, - session_id, - ) - continue - - # Staging: prefill pre-requests staging allocation before forward - if msg[0] == b"STAGING_REQ": - self._handle_staging_req(msg) - continue - - # Prefill acknowledges abort notification - if msg[0] == b"ABORT_ACK": - # TODO(shangming): use this info to implement the deferred release mechanism if needed - ack_aborted_room = int(msg[1].decode("ascii")) - logger.debug(f"Received ABORT_ACK for room {ack_aborted_room}") - 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] + try: + msg = self.server_socket.recv_multipart() + if msg[0] == MooncakeKVManager.AUX_DATA_HEADER: + self._handle_aux_data(msg) + continue + + # Staging: prefill notifies a chunk written to staging buffer + if msg[0] == b"CHUNK_READY": + room = int(msg[1].decode("ascii")) + chunk_idx = int(msg[2].decode("ascii")) + page_start = int(msg[3].decode("ascii")) + num_pages = int(msg[4].decode("ascii")) + session_id = msg[5].decode("ascii") + handler = self._staging_handler + assert ( + handler is not None + ), "CHUNK_READY received before staging handler initialized" + handler.handle_chunk_arrived( + room, + chunk_idx, + page_start, + num_pages, + session_id, ) - arrived_response_num = len( - self.prefill_response_tracker[bootstrap_room] + continue + + # Staging: prefill pre-requests staging allocation before forward + if msg[0] == b"STAGING_REQ": + self._handle_staging_req(msg) + continue + + # Prefill acknowledges abort notification + if msg[0] == b"ABORT_ACK": + # TODO(shangming): use this info to implement the deferred release mechanism if needed + ack_aborted_room = int(msg[1].decode("ascii")) + logger.debug(f"Received ABORT_ACK for room {ack_aborted_room}") + 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", ) - 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) + self.update_status(bootstrap_room, status) + except Exception: + logger.exception("Decode status worker failed") threading.Thread(target=decode_thread).start() self._start_heartbeat_checker_thread() diff --git a/python/sglang/srt/disaggregation/nixl/conn.py b/python/sglang/srt/disaggregation/nixl/conn.py index 04c48ffc80e4..cd84cd2e4e0d 100644 --- a/python/sglang/srt/disaggregation/nixl/conn.py +++ b/python/sglang/srt/disaggregation/nixl/conn.py @@ -597,14 +597,17 @@ def _start_decode_staging_thread(self): def decode_staging_thread(): while True: - msg = self.server_socket.recv_multipart() - if msg[0] == b"STAGING_REQ": - self._handle_staging_req(msg) - continue - logger.warning( - "decode_staging_thread: unexpected message tag %s", - msg[0][:20], - ) + try: + msg = self.server_socket.recv_multipart() + if msg[0] == b"STAGING_REQ": + self._handle_staging_req(msg) + continue + logger.warning( + "decode_staging_thread: unexpected message tag %s", + msg[0][:20], + ) + except Exception: + logger.exception("Decode staging worker failed") threading.Thread(target=decode_staging_thread, daemon=True).start() @@ -2668,69 +2671,74 @@ def _start_bootstrap_thread(self): def bootstrap_thread(): """This thread recvs transfer info from the decode engine""" while True: - waiting_req_bytes = self.server_socket.recv_multipart() - logger.debug( - f"Received multipart with total byte size {sum(len(x) for x in waiting_req_bytes)}" - ) + try: + waiting_req_bytes = self.server_socket.recv_multipart() + logger.debug( + f"Received multipart with total byte size {sum(len(x) for x in waiting_req_bytes)}" + ) - # Staging: decode reports consumption watermark back to prefill - if waiting_req_bytes[0] == b"WATERMARK": - if self.enable_staging: - from sglang.srt.disaggregation.common.staging_handler import ( - handle_watermark_msg, - ) + # Staging: decode reports consumption watermark back to prefill + if waiting_req_bytes[0] == b"WATERMARK": + if self.enable_staging: + from sglang.srt.disaggregation.common.staging_handler import ( + handle_watermark_msg, + ) - handle_watermark_msg(self._staging_ctx, waiting_req_bytes) - continue + handle_watermark_msg(self._staging_ctx, waiting_req_bytes) + continue - # Staging: decode replies with allocated staging offset - if waiting_req_bytes[0] == b"STAGING_RSP": - if self.enable_staging: - from sglang.srt.disaggregation.common.staging_handler import ( - handle_staging_rsp, - ) + # Staging: decode replies with allocated staging offset + if waiting_req_bytes[0] == b"STAGING_RSP": + if self.enable_staging: + from sglang.srt.disaggregation.common.staging_handler import ( + handle_staging_rsp, + ) - handle_staging_rsp(waiting_req_bytes, self.transfer_infos) - continue + handle_staging_rsp(waiting_req_bytes, self.transfer_infos) + continue - if self._handle_abort_notification(waiting_req_bytes): - continue + if self._handle_abort_notification(waiting_req_bytes): + continue - assert ( - waiting_req_bytes[0] == GUARD - ), f"First message should be {GUARD}. Foreign traffic?" - waiting_req_bytes = waiting_req_bytes[1:] - room = waiting_req_bytes[0].decode("ascii") - agent_name = waiting_req_bytes[3].decode("ascii") - if room == "None": - # Register new peer and save KV base pointers. - self._add_remote_peer( - KVArgsRegisterInfo.from_zmq(waiting_req_bytes) + assert ( + waiting_req_bytes[0] == GUARD + ), f"First message should be {GUARD}. Foreign traffic?" + waiting_req_bytes = waiting_req_bytes[1:] + room = waiting_req_bytes[0].decode("ascii") + agent_name = waiting_req_bytes[3].decode("ascii") + if room == "None": + # Register new peer and save KV base pointers. + self._add_remote_peer( + KVArgsRegisterInfo.from_zmq(waiting_req_bytes) + ) + logger.debug(f"Register KVArgs from {agent_name} successfully") + continue + room = int(room) + if room not in self.transfer_infos: + self.transfer_infos[room] = {} + self.transfer_infos[room][agent_name] = TransferInfo.from_zmq( + waiting_req_bytes ) - logger.debug(f"Register KVArgs from {agent_name} successfully") - continue - room = int(room) - if room not in self.transfer_infos: - self.transfer_infos[room] = {} - self.transfer_infos[room][agent_name] = TransferInfo.from_zmq( - waiting_req_bytes - ) - required_dst_info_num = self.transfer_infos[room][ - agent_name - ].required_dst_info_num - logger.debug(f"got info {room=} {agent_name=} {required_dst_info_num=}") - if len(self.transfer_infos[room]) == required_dst_info_num: - self.resolve_kv_replica_factor(self.transfer_infos[room]) - self.req_to_decode_prefix_len[room] = next( - ( - info.decode_prefix_len - for info in self.transfer_infos[room].values() - if info.decode_prefix_len is not None - ), - 0, + required_dst_info_num = self.transfer_infos[room][ + agent_name + ].required_dst_info_num + logger.debug( + f"got info {room=} {agent_name=} {required_dst_info_num=}" ) - logger.debug(f"{room=} is bootstrapped") - self.update_status(room, KVPoll.WaitingForInput) + if len(self.transfer_infos[room]) == required_dst_info_num: + self.resolve_kv_replica_factor(self.transfer_infos[room]) + self.req_to_decode_prefix_len[room] = next( + ( + info.decode_prefix_len + for info in self.transfer_infos[room].values() + if info.decode_prefix_len is not None + ), + 0, + ) + logger.debug(f"{room=} is bootstrapped") + self.update_status(room, KVPoll.WaitingForInput) + except Exception: + logger.exception("Bootstrap worker failed") threading.Thread(target=bootstrap_thread).start() From 48c5e1569d4c5580733213df52738c8f198c7611 Mon Sep 17 00:00:00 2001 From: jambow0320 <1193785411@qq.com> Date: Tue, 18 Aug 2026 01:50:13 +1000 Subject: [PATCH 6/9] align Mori speculative MHA layout with common planner --- python/sglang/srt/disaggregation/mori/conn.py | 12 +++++++++++- 1 file changed, 11 insertions(+), 1 deletion(-) diff --git a/python/sglang/srt/disaggregation/mori/conn.py b/python/sglang/srt/disaggregation/mori/conn.py index 3d5e92ca90de..3b08f95af458 100644 --- a/python/sglang/srt/disaggregation/mori/conn.py +++ b/python/sglang/srt/disaggregation/mori/conn.py @@ -743,8 +743,18 @@ def _get_mha_mem_desc_slices( "Destination KV descriptors do not match prefill pp configuration" ) dst_k_descs = dst_mem_descs[start_layer:end_layer] + if ( + num_local_layers < dst_total_layers + and dst_total_layers % num_local_layers != 0 + ): + # Decode has draft-model KV while Prefill has target-model KV only: + # [K_main..., V_main..., draft_K..., draft_V...]. + multiplier_ratio = dst_total_layers // num_local_layers + dst_v_offset = num_local_layers * multiplier_ratio + else: + dst_v_offset = dst_total_layers dst_v_descs = dst_mem_descs[ - dst_total_layers + start_layer : dst_total_layers + end_layer + dst_v_offset + start_layer : dst_v_offset + end_layer ] return src_k_descs, src_v_descs, dst_k_descs, dst_v_descs, num_local_layers From 12a88dbcfe170a6e80e9c6b588cbbd04bdb0ea2d Mon Sep 17 00:00:00 2001 From: jambow0320 <1193785411@qq.com> Date: Tue, 18 Aug 2026 13:05:28 +1000 Subject: [PATCH 7/9] initialize Mooncake session state before control thread --- python/sglang/srt/disaggregation/mooncake/conn.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py index 3119785b07e1..d21901f17122 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -208,13 +208,13 @@ def __init__( self.enable_staging = envs.SGLANG_DISAGG_STAGING_BUFFER.get() self.enable_trace = server_args.enable_trace if self.disaggregation_mode == DisaggregationMode.PREFILL: - self.start_prefill_thread() self.session_failures = defaultdict(int) self.failed_sessions = set() + self.session_lock = threading.Lock() + self.start_prefill_thread() # Per-room count of chunks not yet transferred; teardown waits for # zero so a deferred chunk is not dropped by an early conclude. self._staging_outstanding = defaultdict(int) - self.session_lock = threading.Lock() # Determine the number of threads to use for kv sender cpu_count = os.cpu_count() transfer_thread_pool_size = ( From 081186c4c88c915e814df5b0be1e3373e473ed3e Mon Sep 17 00:00:00 2001 From: jambow0320 <1193785411@qq.com> Date: Tue, 18 Aug 2026 13:27:32 +1000 Subject: [PATCH 8/9] align Mori same-PP and hybrid MLA layout planning --- python/sglang/srt/disaggregation/mori/conn.py | 20 ++++++++++++++++++- 1 file changed, 19 insertions(+), 1 deletion(-) diff --git a/python/sglang/srt/disaggregation/mori/conn.py b/python/sglang/srt/disaggregation/mori/conn.py index 3b08f95af458..9eb43ad1a5d4 100644 --- a/python/sglang/srt/disaggregation/mori/conn.py +++ b/python/sglang/srt/disaggregation/mori/conn.py @@ -735,6 +735,20 @@ def _get_mha_mem_desc_slices( src_k_descs = src_descs[:num_local_layers] src_v_descs = src_descs[num_local_layers:] + # Both peers expose the same PP-local layout. Their descriptor indices + # are already aligned, so applying the Prefill rank's global layer + # offset would incorrectly index into a local list. + if len(src_descs) == len(dst_mem_descs): + dst_k_descs = dst_mem_descs[:num_local_layers] + dst_v_descs = dst_mem_descs[num_local_layers:] + return ( + src_k_descs, + src_v_descs, + dst_k_descs, + dst_v_descs, + num_local_layers, + ) + start_layer = self.kv_args.prefill_start_layer end_layer = start_layer + num_local_layers dst_total_layers = len(dst_mem_descs) // 2 @@ -763,6 +777,10 @@ def _get_mla_mem_desc_slices( ) -> tuple[List[MemoryDesc], List[MemoryDesc], int]: src_descs = self.kv_mem_descs num_local_layers = len(src_descs) + # Same-PP peers register matching local descriptor lists. + if len(src_descs) == len(dst_mem_descs): + return src_descs, dst_mem_descs, num_local_layers + start_layer = self.kv_args.prefill_start_layer end_layer = start_layer + num_local_layers if end_layer > len(dst_mem_descs): @@ -926,7 +944,7 @@ def send_kvcache( statuses: List[TransferStatus] = [] kv_item_len = self.kv_args.kv_item_lens[0] - if self.is_mla_backend: + if self.is_mla_backend or self.is_hybrid_mla_backend: src_descs, dst_descs, layers_current_pp_stage = ( self._get_mla_mem_desc_slices(peer_info.dst_kv_mem_descs) ) From 8388c351d80d7f64bd70c24d6a31ce2b27ffb77d Mon Sep 17 00:00:00 2001 From: jambow0320 <1193785411@qq.com> Date: Fri, 21 Aug 2026 18:53:55 +1000 Subject: [PATCH 9/9] reconcile defensive alignment with deferred abort handling --- .../srt/disaggregation/mooncake/conn.py | 212 ++++++------------ python/sglang/srt/disaggregation/nixl/conn.py | 135 ++++++----- 2 files changed, 145 insertions(+), 202 deletions(-) diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py index 9b2b0a21255a..2e8451762fbd 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -2040,108 +2040,71 @@ def bootstrap_thread(): f"Received abort notification for room {room_to_be_aborted}, " f"ignoring (already completed or unknown)" ) - - handle_watermark_msg(self._staging_ctx, waiting_req_bytes) - continue - # Staging: decode replies with allocated staging offset - if room == "STAGING_RSP": - from sglang.srt.disaggregation.common.staging_handler import ( - handle_staging_rsp, + # Send ACK back to decode endpoint + try: + na = NetworkAddress(decode_ip, decode_port) + self._send_multipart_locked( + na.to_tcp(), + [ + b"ABORT_ACK", + str(room_to_be_aborted).encode("ascii"), + ], + is_ipv6=na.is_ipv6, ) - - handle_staging_rsp(waiting_req_bytes, self.transfer_infos) - continue - # Decode-side abort notification: mark room as failed and ACK - if room == "ABORT": - room_to_be_aborted = int(waiting_req_bytes[1].decode("ascii")) - decode_ip = waiting_req_bytes[2].decode("ascii") - decode_port = int(waiting_req_bytes[3].decode("ascii")) - # No need to abort the room if it has already succeeded - if ( - room_to_be_aborted in self.request_status - and self.check_status(room_to_be_aborted) != KVPoll.Success - ): - self.update_status(room_to_be_aborted, KVPoll.Failed) - logger.debug( - f"Received abort notification for room {room_to_be_aborted}, " - f"marked as Failed" - ) - else: - logger.debug( - f"Received abort notification for room {room_to_be_aborted}, " - f"ignoring (already completed or unknown)" - ) - # Send ACK back to decode endpoint - try: - na = NetworkAddress(decode_ip, decode_port) - self._send_multipart_locked( - na.to_tcp(), - [ - b"ABORT_ACK", - str(room_to_be_aborted).encode("ascii"), - ], - is_ipv6=na.is_ipv6, - ) - logger.debug( - f"Sent ABORT_ACK for room {room_to_be_aborted} to " - f"{decode_ip}:{decode_port}" - ) - except Exception as e: - logger.debug( - f"Failed to send ABORT_ACK for room {room_to_be_aborted}: {e}" - ) - continue - mooncake_session_id = waiting_req_bytes[3].decode("ascii") - if room == "None": - decode_kv_args = KVArgsRegisterInfo.from_zmq(waiting_req_bytes) - decode_kv_args.requires_dcp_relayout = ( - self.requires_dcp_relayout( - decode_kv_args.dst_dcp_size, - decode_kv_args.dst_dcp_rank, - ) + logger.debug( + f"Sent ABORT_ACK for room {room_to_be_aborted} to " + f"{decode_ip}:{decode_port}" ) - if decode_kv_args.requires_dcp_relayout: - decode_kv_args.dcp_token_item_lens = ( - self.prepare_dcp_token_item_lens( - [decode_kv_args.dst_kv_item_len] - * len(self.kv_args.kv_item_lens) - ) - ) - self.decode_kv_args_table[mooncake_session_id] = decode_kv_args - with self.session_lock: - if mooncake_session_id in self.failed_sessions: - self.failed_sessions.remove(mooncake_session_id) - if mooncake_session_id in self.session_failures: - del self.session_failures[mooncake_session_id] + except Exception as e: logger.debug( - f"Register KVArgs from {mooncake_session_id} successfully" + f"Failed to send ABORT_ACK for room {room_to_be_aborted}: {e}" ) - continue - else: - required_dst_info_num = int( - waiting_req_bytes[7].decode("ascii") + continue + mooncake_session_id = waiting_req_bytes[3].decode("ascii") + if room == "None": + decode_kv_args = KVArgsRegisterInfo.from_zmq(waiting_req_bytes) + decode_kv_args.requires_dcp_relayout = self.requires_dcp_relayout( + decode_kv_args.dst_dcp_size, + decode_kv_args.dst_dcp_rank, + ) + if decode_kv_args.requires_dcp_relayout: + decode_kv_args.dcp_token_item_lens = ( + self.prepare_dcp_token_item_lens( + [decode_kv_args.dst_kv_item_len] + * len(self.kv_args.kv_item_lens) + ) ) - room = int(room) - if room not in self.transfer_infos: - self.transfer_infos[room] = {} + self.decode_kv_args_table[mooncake_session_id] = decode_kv_args + with self.session_lock: + if mooncake_session_id in self.failed_sessions: + self.failed_sessions.remove(mooncake_session_id) + if mooncake_session_id in self.session_failures: + del self.session_failures[mooncake_session_id] + logger.debug( + f"Register KVArgs from {mooncake_session_id} successfully" + ) + continue + else: + required_dst_info_num = int(waiting_req_bytes[7].decode("ascii")) + room = int(room) + if room not in self.transfer_infos: + self.transfer_infos[room] = {} - self.transfer_infos[room][mooncake_session_id] = ( - TransferInfo.from_zmq(waiting_req_bytes) + self.transfer_infos[room][mooncake_session_id] = ( + TransferInfo.from_zmq(waiting_req_bytes) + ) + # NOTE: after bootstrapping we can mark the req as waiting for input + if len(self.transfer_infos[room]) == required_dst_info_num: + self.resolve_kv_replica_factor(self.transfer_infos[room]) + self.req_to_decode_prefix_len[room] = next( + ( + info.decode_prefix_len + for info in self.transfer_infos[room].values() + if info.decode_prefix_len is not None + ), + 0, ) - # NOTE: after bootstrapping we can mark the req as waiting for input - if len(self.transfer_infos[room]) == required_dst_info_num: - self.resolve_kv_replica_factor(self.transfer_infos[room]) - self.req_to_decode_prefix_len[room] = next( - ( - info.decode_prefix_len - for info in self.transfer_infos[room].values() - if info.decode_prefix_len is not None - ), - 0, - ) - self.update_status(room, KVPoll.WaitingForInput) - except Exception: - logger.exception("Bootstrap worker failed") + self.update_status(room, KVPoll.WaitingForInput) threading.Thread(target=bootstrap_thread).start() @@ -2201,52 +2164,21 @@ def decode_thread(): expected_response_num = ( self.required_prefill_response_num_table[bootstrap_room] ) - continue - - # Staging: prefill pre-requests staging allocation before forward - if msg[0] == b"STAGING_REQ": - self._handle_staging_req(msg) - continue - - # Prefill acknowledges abort notification - if msg[0] == b"ABORT_ACK": - # TODO(shangming): use this info to implement the deferred release mechanism if needed - ack_aborted_room = int(msg[1].decode("ascii")) - logger.debug(f"Received ABORT_ACK for room {ack_aborted_room}") - 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", + arrived_response_num = len( + self.prefill_response_tracker[bootstrap_room] ) - self.update_status(bootstrap_room, status) - except Exception: - logger.exception("Decode status worker failed") + 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) threading.Thread(target=decode_thread).start() self._start_heartbeat_checker_thread() diff --git a/python/sglang/srt/disaggregation/nixl/conn.py b/python/sglang/srt/disaggregation/nixl/conn.py index 65a79a4c8249..9d45831130ad 100644 --- a/python/sglang/srt/disaggregation/nixl/conn.py +++ b/python/sglang/srt/disaggregation/nixl/conn.py @@ -1129,6 +1129,14 @@ def transfer_worker(self, queue: FastQueue, staging_buffer=None): room = kv_chunk.room handles: List[Any] = [] try: + if room not in self.request_status: + logger.debug( + "Skipping chunk for room %s because it has been cleared", + room, + ) + self._staging_outstanding.pop(room, None) + continue + # Counted at dequeue, before the status check, so # `outstanding == 0` means nothing is dequeued or in flight -- # the predicate the abort ack relies on. The flag survives @@ -1144,7 +1152,15 @@ def transfer_worker(self, queue: FastQueue, staging_buffer=None): self._maybe_ack_drained_abort(room) continue - assert room in self.transfer_infos + room_transfer_infos = self.transfer_infos.get(room) + if room_transfer_infos is None: + logger.debug( + "Skipping chunk for room %s because its transfer metadata " + "has been cleared", + room, + ) + self._staging_outstanding.pop(room, None) + continue # Lazily build a per-worker staging strategy bound to this # worker's private staging buffer (matches mooncake). @@ -1157,7 +1173,7 @@ def transfer_worker(self, queue: FastQueue, staging_buffer=None): self.update_status(room, KVPoll.Transferring) - reqs_to_be_processed = list(self.transfer_infos[room].values()) + reqs_to_be_processed = list(room_transfer_infos.values()) # Set when staging allocation/watermark is not yet ready and # the chunk has been re-enqueued. We then break out of the @@ -2697,74 +2713,69 @@ def _start_bootstrap_thread(self): def bootstrap_thread(): """This thread recvs transfer info from the decode engine""" while True: - try: - waiting_req_bytes = self.server_socket.recv_multipart() - logger.debug( - f"Received multipart with total byte size {sum(len(x) for x in waiting_req_bytes)}" - ) + waiting_req_bytes = self.server_socket.recv_multipart() + logger.debug( + f"Received multipart with total byte size {sum(len(x) for x in waiting_req_bytes)}" + ) - # Staging: decode reports consumption watermark back to prefill - if waiting_req_bytes[0] == b"WATERMARK": - if self.enable_staging: - from sglang.srt.disaggregation.common.staging_handler import ( - handle_watermark_msg, - ) + # Staging: decode reports consumption watermark back to prefill + if waiting_req_bytes[0] == b"WATERMARK": + if self.enable_staging: + from sglang.srt.disaggregation.common.staging_handler import ( + handle_watermark_msg, + ) - handle_watermark_msg(self._staging_ctx, waiting_req_bytes) - continue + handle_watermark_msg(self._staging_ctx, waiting_req_bytes) + continue - # Staging: decode replies with allocated staging offset - if waiting_req_bytes[0] == b"STAGING_RSP": - if self.enable_staging: - from sglang.srt.disaggregation.common.staging_handler import ( - handle_staging_rsp, - ) + # Staging: decode replies with allocated staging offset + if waiting_req_bytes[0] == b"STAGING_RSP": + if self.enable_staging: + from sglang.srt.disaggregation.common.staging_handler import ( + handle_staging_rsp, + ) - handle_staging_rsp(waiting_req_bytes, self.transfer_infos) - continue + handle_staging_rsp(waiting_req_bytes, self.transfer_infos) + continue - if self._handle_abort_notification(waiting_req_bytes): - continue + if self._handle_abort_notification(waiting_req_bytes): + continue - assert ( - waiting_req_bytes[0] == GUARD - ), f"First message should be {GUARD}. Foreign traffic?" - waiting_req_bytes = waiting_req_bytes[1:] - room = waiting_req_bytes[0].decode("ascii") - agent_name = waiting_req_bytes[3].decode("ascii") - if room == "None": - # Register new peer and save KV base pointers. - self._add_remote_peer( - KVArgsRegisterInfo.from_zmq(waiting_req_bytes) - ) - logger.debug(f"Register KVArgs from {agent_name} successfully") - continue - room = int(room) - if room not in self.transfer_infos: - self.transfer_infos[room] = {} - self.transfer_infos[room][agent_name] = TransferInfo.from_zmq( - waiting_req_bytes + assert ( + waiting_req_bytes[0] == GUARD + ), f"First message should be {GUARD}. Foreign traffic?" + waiting_req_bytes = waiting_req_bytes[1:] + room = waiting_req_bytes[0].decode("ascii") + agent_name = waiting_req_bytes[3].decode("ascii") + if room == "None": + # Register new peer and save KV base pointers. + self._add_remote_peer( + KVArgsRegisterInfo.from_zmq(waiting_req_bytes) ) - required_dst_info_num = self.transfer_infos[room][ - agent_name - ].required_dst_info_num - logger.debug( - f"got info {room=} {agent_name=} {required_dst_info_num=}" + logger.debug(f"Register KVArgs from {agent_name} successfully") + continue + room = int(room) + if room not in self.transfer_infos: + self.transfer_infos[room] = {} + self.transfer_infos[room][agent_name] = TransferInfo.from_zmq( + waiting_req_bytes + ) + required_dst_info_num = self.transfer_infos[room][ + agent_name + ].required_dst_info_num + logger.debug(f"got info {room=} {agent_name=} {required_dst_info_num=}") + if len(self.transfer_infos[room]) == required_dst_info_num: + self.resolve_kv_replica_factor(self.transfer_infos[room]) + self.req_to_decode_prefix_len[room] = next( + ( + info.decode_prefix_len + for info in self.transfer_infos[room].values() + if info.decode_prefix_len is not None + ), + 0, ) - if len(self.transfer_infos[room]) == required_dst_info_num: - self.resolve_kv_replica_factor(self.transfer_infos[room]) - self.req_to_decode_prefix_len[room] = next( - ( - info.decode_prefix_len - for info in self.transfer_infos[room].values() - if info.decode_prefix_len is not None - ), - 0, - ) - logger.debug(f"{room=} is bootstrapped") - self.update_status(room, KVPoll.WaitingForInput) - except Exception: - logger.exception("Bootstrap worker failed") + logger.debug(f"{room=} is bootstrapped") + self.update_status(room, KVPoll.WaitingForInput) threading.Thread(target=bootstrap_thread).start()