[Bugfix] Clean up failed NIXL sender state (#27011)
This commit is contained in:
@@ -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()
|
||||
Reference in New Issue
Block a user