diff --git a/python/sglang/srt/disaggregation/nixl/conn.py b/python/sglang/srt/disaggregation/nixl/conn.py index 0fbdebde9..2554439e9 100644 --- a/python/sglang/srt/disaggregation/nixl/conn.py +++ b/python/sglang/srt/disaggregation/nixl/conn.py @@ -1944,12 +1944,36 @@ class NixlKVSender(CommonKVSender): ) return status + def clear(self) -> None: + super().clear() + if ( + getattr(self.kv_mgr, "enable_staging", False) + and getattr(self.kv_mgr, "_staging_ctx", None) is not None + ): + self.kv_mgr._staging_ctx.prefetched_rooms.discard(self.bootstrap_room) + self.kv_mgr._staging_ctx.prefetch_requested = { + key + for key in self.kv_mgr._staging_ctx.prefetch_requested + if key[0] != self.bootstrap_room + } + def failure_exception(self): + exc = self.kv_mgr.exceptions.pop(self.bootstrap_room, None) + with self.kv_mgr.failure_lock: + failure_reason = self.kv_mgr.failure_records.pop(self.bootstrap_room, None) + + if self.conclude_state is None: + self.conclude_state = KVPoll.Failed + self._send_failed = True + + self.clear() + if self._send_error is not None: raise self._send_error - exc = self.kv_mgr.exceptions.pop(self.bootstrap_room, None) if exc is not None: raise exc + if failure_reason is not None: + raise RuntimeError(failure_reason) raise RuntimeError("NIXL KVSender Exception") diff --git a/test/registered/unit/disaggregation/test_nixl_sender_failure_cleanup.py b/test/registered/unit/disaggregation/test_nixl_sender_failure_cleanup.py new file mode 100644 index 000000000..81d3bb423 --- /dev/null +++ b/test/registered/unit/disaggregation/test_nixl_sender_failure_cleanup.py @@ -0,0 +1,54 @@ +import threading +import unittest +from types import SimpleNamespace + +from sglang.srt.disaggregation.base.conn import KVPoll +from sglang.srt.disaggregation.nixl.conn import NixlKVSender +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=2, suite="base-a-test-cpu") + + +class TestNixlSenderFailureCleanup(unittest.TestCase): + def test_failure_exception_cleans_room_state_before_raising(self): + room = 7 + expected_exc = RuntimeError("transfer failed") + sender = NixlKVSender.__new__(NixlKVSender) + sender.bootstrap_room = room + sender.conclude_state = None + sender._send_failed = False + sender._send_error = None + staging_ctx = SimpleNamespace( + prefetched_rooms={room, 8}, + prefetch_requested={(room, 0, "session-a"), (8, 0, "session-b")}, + ) + sender.kv_mgr = SimpleNamespace( + enable_staging=True, + _staging_ctx=staging_ctx, + request_status={room: object()}, + req_to_decode_prefix_len={room: 3}, + transfer_infos={room: object()}, + exceptions={room: expected_exc}, + failure_records={room: "transfer failed"}, + failure_lock=threading.Lock(), + ) + + with self.assertRaises(RuntimeError) as cm: + sender.failure_exception() + + self.assertIs(cm.exception, expected_exc) + self.assertTrue(sender._send_failed) + self.assertEqual(sender.conclude_state, KVPoll.Failed) + self.assertNotIn(room, sender.kv_mgr.request_status) + self.assertNotIn(room, sender.kv_mgr.req_to_decode_prefix_len) + self.assertNotIn(room, sender.kv_mgr.transfer_infos) + self.assertNotIn(room, sender.kv_mgr.exceptions) + self.assertNotIn(room, sender.kv_mgr.failure_records) + self.assertNotIn(room, staging_ctx.prefetched_rooms) + self.assertNotIn((room, 0, "session-a"), staging_ctx.prefetch_requested) + self.assertIn(8, staging_ctx.prefetched_rooms) + self.assertIn((8, 0, "session-b"), staging_ctx.prefetch_requested) + + +if __name__ == "__main__": + unittest.main()