[Bugfix] Clean up failed NIXL sender state (#27011)

This commit is contained in:
CrazyCoder
2026-06-03 12:15:15 +08:00
committed by GitHub
parent 3e681d7fff
commit c3aaafc5f2
2 changed files with 79 additions and 1 deletions
+25 -1
View File
@@ -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")
@@ -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()