Files
sglang/test/registered/unit/disaggregation/test_mooncake_abort_protocol.py
leavelet 9a40e42384 fix(disagg/mooncake): notify prefill on decode-side abort
Port the decode->prefill half of upstream sgl-project/sglang #27372 (the worker
skip-Failed guard half landed in cb1a03f0a). In multi-prefill/multi-decode PD, a
decode-initiated abort frees the decode's KV pages back to the allocator, but the
prefill never learns of it (request_status is per-process), so the prefill keeps
RDMA-writing the remaining chunks into pages the decode may have reallocated to a
different live request -> KV corruption. The already-ported worker guard is a no-op
here because nothing sets the prefill room to Failed.

Now the decode receiver sends a 4-field b"ABORT" notification to every prefill peer
(on abort() and on the poll() WaitingForInput timeout, at most once); the prefill
marks the room Failed (so the worker guard fires) and replies b"ABORT_ACK". The
ABORT branch is handled before the unconditional waiting_req_bytes[3] decode in
bootstrap_thread (a 4-field ABORT would otherwise crash the else branch), and
ABORT_ACK before the 3-tuple unpack in the decode thread.

De-entangled from the staging/tracing infra this branch lacks. Component-tested
(test_mooncake_abort_protocol.py: format, send-once guard, abort wiring, ZMQ
round-trip). Cross-process corruption-window closure to be verified by the
multi-P/D PD harness.

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-06-07 09:51:04 +00:00

139 lines
4.6 KiB
Python

"""Component tests for the decode->prefill abort protocol (port of upstream #27372).
These cover the pieces that do NOT need a live PD transfer:
- the decode side builds the exact 4-field b"ABORT" message and sends it at most once,
- MooncakeKVReceiver.abort() invokes the notification and marks the room Failed,
- the ZMQ wire format round-trips so the prefill handler's waiting_req_bytes[1..3]
indices are always present (the branch-ordering trap: a 4-field ABORT must not
IndexError before the room=="ABORT" check).
The actual cross-process corruption-window closure is verified by the multi-P/D PD
harness, not here.
"""
import unittest
from sglang.srt.disaggregation.base.conn import KVPoll
from sglang.srt.disaggregation.mooncake.conn import MooncakeKVReceiver
class _DummyLock:
def __enter__(self):
return self
def __exit__(self, *a):
return False
class _CapturingSock:
def __init__(self):
self.sent = []
def send_multipart(self, parts):
self.sent.append(parts)
class _FakeMgr:
def __init__(self):
self.local_ip = "10.0.0.1"
self.rank_port = 12345
self.status_updates = []
def record_failure(self, room, reason):
pass
def update_status(self, room, status):
self.status_updates.append((room, status))
def _make_receiver(bootstrap_infos):
r = MooncakeKVReceiver.__new__(MooncakeKVReceiver) # bypass heavy __init__
r.abort_notified = False
r.bootstrap_room = 777
r.bootstrap_infos = bootstrap_infos
r.kv_mgr = _FakeMgr()
r.conclude_state = None
return r
class TestSendAbortNotification(unittest.TestCase):
def test_message_format_is_four_fields(self):
sock = _CapturingSock()
r = _make_receiver([{"peer": 1}])
r._connect_to_bootstrap_server = lambda info: (sock, _DummyLock())
r._send_abort_notification()
self.assertEqual(len(sock.sent), 1)
parts = sock.sent[0]
# [ABORT, room, ip, port] — exactly 4 fields so the prefill handler's
# waiting_req_bytes[3] (port) access never IndexErrors.
self.assertEqual(len(parts), 4)
self.assertEqual(parts[0], b"ABORT")
self.assertEqual(parts[1], b"777")
self.assertEqual(parts[2], b"10.0.0.1")
self.assertEqual(parts[3], b"12345")
def test_sent_at_most_once(self):
sock = _CapturingSock()
r = _make_receiver([{"peer": 1}])
r._connect_to_bootstrap_server = lambda info: (sock, _DummyLock())
r._send_abort_notification()
r._send_abort_notification()
self.assertEqual(len(sock.sent), 1)
self.assertTrue(r.abort_notified)
def test_no_bootstrap_infos_is_noop(self):
r = _make_receiver(None)
r._send_abort_notification() # must not raise
self.assertFalse(r.abort_notified)
def test_sends_to_every_peer(self):
socks = [_CapturingSock(), _CapturingSock()]
infos = [{"peer": 0}, {"peer": 1}]
r = _make_receiver(infos)
r._connect_to_bootstrap_server = lambda info: (socks[info["peer"]], _DummyLock())
r._send_abort_notification()
self.assertEqual(len(socks[0].sent), 1)
self.assertEqual(len(socks[1].sent), 1)
class TestAbortWiring(unittest.TestCase):
def test_abort_notifies_and_marks_failed(self):
r = _make_receiver([{"peer": 1}])
called = []
r._send_abort_notification = lambda: called.append(True)
r.abort()
self.assertTrue(called)
self.assertEqual(r.conclude_state, KVPoll.Failed)
self.assertIn((777, KVPoll.Failed), r.kv_mgr.status_updates)
class TestZmqWireFormat(unittest.TestCase):
def test_abort_and_ack_roundtrip(self):
import zmq
ctx = zmq.Context.instance()
pull = ctx.socket(zmq.PULL)
pull.bind("tcp://127.0.0.1:*")
ep = pull.getsockopt_string(zmq.LAST_ENDPOINT)
push = ctx.socket(zmq.PUSH)
push.connect(ep)
try:
push.send_multipart([b"ABORT", b"777", b"10.0.0.1", b"12345"])
parts = pull.recv_multipart()
self.assertEqual(parts[0], b"ABORT")
self.assertEqual(int(parts[1].decode()), 777)
self.assertEqual(parts[2].decode(), "10.0.0.1")
self.assertEqual(int(parts[3].decode()), 12345)
push.send_multipart([b"ABORT_ACK", b"777"])
ack = pull.recv_multipart()
self.assertEqual(ack[0], b"ABORT_ACK")
self.assertEqual(int(ack[1].decode()), 777)
finally:
push.close()
pull.close()
if __name__ == "__main__":
unittest.main()