From 22cf7d2b42fe34bfb331c8a737715b5ae823df91 Mon Sep 17 00:00:00 2001 From: cctry Date: Tue, 5 May 2026 17:23:42 -0700 Subject: [PATCH] [Fix] Handle nixlRemoteDisconnectError in NixlKVSender (#24296) --- python/sglang/srt/disaggregation/nixl/conn.py | 71 +++++++++++++++---- 1 file changed, 58 insertions(+), 13 deletions(-) diff --git a/python/sglang/srt/disaggregation/nixl/conn.py b/python/sglang/srt/disaggregation/nixl/conn.py index 471ab416f..5c11b3efb 100644 --- a/python/sglang/srt/disaggregation/nixl/conn.py +++ b/python/sglang/srt/disaggregation/nixl/conn.py @@ -28,6 +28,21 @@ from sglang.srt.disaggregation.utils import ( from sglang.srt.environ import envs from sglang.srt.server_args import ServerArgs +try: + from nixl._bindings import ( + nixlBackendError, + nixlCancelledError, + nixlRemoteDisconnectError, + ) + + _NIXL_TRANSPORT_ERRORS = ( + nixlRemoteDisconnectError, + nixlBackendError, + nixlCancelledError, + ) +except ImportError: + _NIXL_TRANSPORT_ERRORS = (RuntimeError,) + logger = logging.getLogger(__name__) GUARD = "NixlMsgGuard".encode("ascii") @@ -1098,6 +1113,8 @@ class NixlKVSender(CommonKVSender): self.xfer_handles = [] self.has_sent = False self.chunk_id = 0 + self._send_failed = False + self._send_error: Optional[Exception] = None def pop_decode_prefix_len(self) -> int: return self.kv_mgr.req_to_decode_prefix_len.pop(self.bootstrap_room, 0) @@ -1110,6 +1127,9 @@ class NixlKVSender(CommonKVSender): kv_indices: npt.NDArray[np.int32], state_indices: Optional[List[int]] = None, ): + if self._send_failed: + return + index_slice = slice(self.curr_idx, self.curr_idx + len(kv_indices)) self.curr_idx += len(kv_indices) is_last = self.curr_idx == self.num_kv_indices @@ -1128,15 +1148,24 @@ class NixlKVSender(CommonKVSender): self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Success) return - new_xfer_handles = self.kv_mgr.add_transfer_request( - self.bootstrap_room, - kv_indices, - index_slice, - is_last, - self.chunk_id, - self.aux_index, - state_indices, - ) + try: + new_xfer_handles = self.kv_mgr.add_transfer_request( + self.bootstrap_room, + kv_indices, + index_slice, + is_last, + self.chunk_id, + self.aux_index, + state_indices, + ) + except _NIXL_TRANSPORT_ERRORS as e: + logger.warning( + f"KVSender transfer request failed for room {self.bootstrap_room}: {e}" + ) + self._send_failed = True + self._send_error = e + return + self.xfer_handles.extend(new_xfer_handles) self.chunk_id += 1 if is_last: @@ -1144,16 +1173,32 @@ class NixlKVSender(CommonKVSender): del self.kv_mgr.request_status[self.bootstrap_room] def poll(self) -> KVPoll: + if self._send_failed: + return KVPoll.Failed # type: ignore if not self.has_sent: return self.kv_mgr.check_status(self.bootstrap_room) - states = [self.kv_mgr.agent.check_xfer_state(x) for x in self.xfer_handles] - if all([x == "DONE" for x in states]): + try: + states = [self.kv_mgr.agent.check_xfer_state(x) for x in self.xfer_handles] + except _NIXL_TRANSPORT_ERRORS as e: + logger.warning( + f"KVSender check_xfer_state failed for room {self.bootstrap_room}: {e}" + ) + self._send_failed = True + self._send_error = e + return KVPoll.Failed # type: ignore + if all(x == "DONE" for x in states): return KVPoll.Success # type: ignore - if any([x == "ERR" for x in states]): - raise Exception("KVSender transfer encountered an error.") + if any(x == "ERR" for x in states): + self._send_failed = True + self._send_error = RuntimeError( + f"NIXL transfer error for room {self.bootstrap_room}" + ) + return KVPoll.Failed # type: ignore return KVPoll.WaitingForInput # type: ignore def failure_exception(self): + if self._send_error is not None: + raise self._send_error raise RuntimeError("NIXL KVSender Exception")