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>
139 lines
4.6 KiB
Python
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()
|