[PD] Handle NIXL abort notifications (#30352)

This commit is contained in:
Yichao Cheng
2026-07-13 05:05:25 -07:00
committed by GitHub
parent eb31b5310c
commit a74bee2261
2 changed files with 222 additions and 0 deletions
@@ -631,6 +631,12 @@ class NixlKVManager(CommonKVManager):
def check_status(self, bootstrap_room: int):
return self.request_status.get(bootstrap_room, KVPoll.WaitingForInput)
def update_status(self, bootstrap_room: int, status: KVPoll):
# Keep Failed sticky until the sender clears the room.
if self.request_status.get(bootstrap_room) == KVPoll.Failed:
return
super().update_status(bootstrap_room, status)
def _prep_equal_tp_dlist(
self,
peer_name: str,
@@ -2316,6 +2322,40 @@ class NixlKVManager(CommonKVManager):
return False
return self.transfer_statuses[room].is_done()
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}, "
f"marked as Failed"
)
else:
logger.debug(
f"Received abort notification for room {room_to_be_aborted}, "
f"ignoring (already completed or unknown)"
)
# TODO: Define real ACK/deferred-release semantics if decode-side buffer
# release needs to wait for prefill-side NIXL transfer quiescence.
return True
def _start_bootstrap_thread(self):
def bootstrap_thread():
"""This thread recvs transfer info from the decode engine"""
@@ -2345,6 +2385,9 @@ class NixlKVManager(CommonKVManager):
handle_staging_rsp(waiting_req_bytes, self.transfer_infos)
continue
if self._handle_abort_notification(waiting_req_bytes):
continue
assert (
waiting_req_bytes[0] == GUARD
), f"First message should be {GUARD}. Foreign traffic?"
@@ -330,6 +330,185 @@ class TestNixlKVSenderChunkPolicy(CustomTestCase):
self.assertTrue(sender.should_send_kv_chunk(3, last_chunk=False))
class TestNixlAbortHandling(CustomTestCase):
def _make_manager(self, request_status=None):
mgr = object.__new__(NixlKVManager)
mgr.request_status = dict(request_status or {})
mgr._connect = MagicMock()
mgr.failure_lock = threading.Lock()
mgr.failure_records = {}
return mgr
def test_given_known_incomplete_room_when_abort_arrives_then_room_fails_without_ack(
self,
):
mgr = self._make_manager({11: KVPoll.WaitingForInput})
handled = mgr._handle_abort_notification(
[b"ABORT", b"11", b"127.0.0.1", b"5555"]
)
self.assertTrue(handled)
self.assertEqual(mgr.request_status[11], KVPoll.Failed)
self.assertEqual(
mgr.failure_records[11],
"Aborted by decode-side abort notification.",
)
mgr._connect.assert_not_called()
def test_given_successful_room_when_abort_arrives_then_status_is_preserved(self):
mgr = self._make_manager({12: KVPoll.Success})
handled = mgr._handle_abort_notification(
[b"ABORT", b"12", b"127.0.0.1", b"5556"]
)
self.assertTrue(handled)
self.assertEqual(mgr.request_status[12], KVPoll.Success)
self.assertEqual(mgr.failure_records, {})
mgr._connect.assert_not_called()
def test_given_unknown_room_when_abort_arrives_then_status_remains_absent(self):
mgr = self._make_manager()
handled = mgr._handle_abort_notification(
[b"ABORT", b"14", b"127.0.0.1", b"5557"]
)
self.assertTrue(handled)
self.assertNotIn(14, mgr.request_status)
self.assertEqual(mgr.failure_records, {})
mgr._connect.assert_not_called()
def test_given_malformed_abort_when_handled_then_no_exception_or_ack(self):
mgr = self._make_manager({13: KVPoll.WaitingForInput})
handled = mgr._handle_abort_notification(
[b"ABORT", b"invalid-room", b"127.0.0.1", b"5558"]
)
self.assertTrue(handled)
self.assertEqual(mgr.request_status[13], KVPoll.WaitingForInput)
self.assertEqual(mgr.failure_records, {})
mgr._connect.assert_not_called()
class TestNixlUpdateStatus(CustomTestCase):
def _make_manager(self, request_status):
mgr = object.__new__(NixlKVManager)
mgr.request_status = dict(request_status)
return mgr
def test_given_failed_room_when_status_is_promoted_then_failed_is_preserved(self):
for status in (KVPoll.Transferring, KVPoll.Success):
with self.subTest(status=status):
mgr = self._make_manager({17: KVPoll.Failed})
mgr.update_status(17, status)
self.assertEqual(mgr.request_status[17], KVPoll.Failed)
def test_given_missing_room_when_failed_update_arrives_then_room_is_not_resurrected(
self,
):
mgr = self._make_manager({})
mgr.update_status(18, KVPoll.Failed)
self.assertNotIn(18, mgr.request_status)
class TestNixlTransferWorker(CustomTestCase):
def _make_manager(self, room):
mgr = object.__new__(NixlKVManager)
mgr.request_status = {room: KVPoll.WaitingForInput}
mgr.transfer_infos = {
room: {
"agent": TransferInfo(
room=room,
endpoint="127.0.0.1",
dst_port=5555,
agent_name="agent",
dst_kv_indices=np.array([2], dtype=np.int32),
dst_aux_index=0,
required_dst_info_num=1,
dst_state_indices=[],
)
}
}
mgr.decode_kv_args_table = {
"agent": SimpleNamespace(
decode_tp_size=1,
dst_kv_ptrs=[0],
dst_aux_ptrs=[0],
gpu_id=0,
staging=None,
kv_xfer_segments=None,
dst_homogeneous_mem_kind="VRAM",
)
}
mgr.req_to_decode_prefix_len = {room: 4}
mgr.enable_staging = False
mgr._staging_ctx = None
mgr.is_mla_backend = False
mgr.attn_tp_size = 1
mgr.kv_args = SimpleNamespace(engine_rank=0)
mgr.exceptions = {}
mgr.failure_lock = threading.Lock()
mgr.failure_records = {}
def check_xfer_state(_handle):
mgr.update_status(room, KVPoll.Failed)
return "DONE"
mgr.agent = SimpleNamespace(check_xfer_state=check_xfer_state)
return mgr
def _make_chunk(self, room, prefill_kv_indices, is_last_chunk):
return TransferKVChunk(
room=room,
prefill_kv_indices=np.array(prefill_kv_indices, dtype=np.int32),
index_slice=slice(0, len(prefill_kv_indices)),
is_last_chunk=is_last_chunk,
chunk_id=0,
prefill_aux_index=0 if is_last_chunk else None,
state_indices=None,
)
def _run_worker_once(self, mgr, chunk):
queue = SimpleNamespace(get=MagicMock(side_effect=[chunk, SystemExit()]))
with self.assertRaises(SystemExit):
mgr.transfer_worker(queue)
def test_given_last_chunk_aborts_mid_transfer_when_worker_finishes_then_failed_status_is_preserved(
self,
):
room = 21
mgr = self._make_manager(room)
mgr.send_aux = MagicMock(return_value="aux_handle")
chunk = self._make_chunk(room, [], is_last_chunk=True)
self._run_worker_once(mgr, chunk)
self.assertEqual(mgr.request_status[room], KVPoll.Failed)
self.assertNotIn(room, mgr.transfer_infos)
self.assertNotIn(room, mgr.req_to_decode_prefix_len)
def test_given_non_last_chunk_aborts_mid_transfer_when_worker_finishes_then_failed_status_is_preserved(
self,
):
room = 22
mgr = self._make_manager(room)
mgr.send_kvcache = MagicMock(return_value="kv_handle")
chunk = self._make_chunk(room, [1], is_last_chunk=False)
self._run_worker_once(mgr, chunk)
self.assertEqual(mgr.request_status[room], KVPoll.Failed)
self.assertIn(room, mgr.transfer_infos)
self.assertIn(room, mgr.req_to_decode_prefix_len)
class TestNixlNotifications(CustomTestCase):
def _make_manager(self, messages, required=None):
mgr = object.__new__(NixlKVManager)