"""Unit tests for A2 core: PerLayerTransferContext (per-request overlapped transfer). Key correctness properties: - submit_layer waits the layer's CUDA write event BEFORE issuing the RDMA (so the transfer never reads a layer before its write kernel finished); - idempotent per layer; skips layers with no owned blocks; - a submit failure marks the context failed and stops further submits; - finish waits all accumulated batch_ids and returns 0 / non-zero correctly. Engine and CUDA events are mocked — no GPU/RDMA. """ import unittest from sglang.srt.disaggregation.cp_per_layer_transfer import PerLayerTransferContext class _FakeEvent: def __init__(self): self.synced = False def synchronize(self): self.synced = True class _FakeEngine: def __init__(self, submit_rets=None, wait_ret=0): self.submits = [] self.waits = [] self._rets = submit_rets self._wait = wait_ret self._n = 0 self.submit_order = [] # (layer-derived addr, whether event was synced first) def batch_transfer_async_submit(self, sid, src, dst, lens): self.submits.append((sid, list(src), list(dst), list(lens))) if self._rets is not None: r = self._rets[self._n] self._n += 1 return r return 100 + len(self.submits) def wait_batch_transfers(self, ids): self.waits.append(list(ids)) return self._wait def _blocks(layer_id): base = 1000 * layer_id return ([base], [base + 5000], [64]) def _ctx(engine, blocks=_blocks, num_layers=8): return PerLayerTransferContext(engine, "sess", blocks, num_layers=num_layers) class TestPerLayerTransferContext(unittest.TestCase): def test_submit_waits_event_then_submits(self): eng = _FakeEngine() ev = _FakeEvent() _ctx(eng).submit_layer(3, ev) self.assertTrue(ev.synced) # write event waited BEFORE the RDMA submit self.assertEqual(eng.submits, [("sess", [3000], [8000], [64])]) def test_idempotent_per_layer(self): eng = _FakeEngine() ctx = _ctx(eng) ctx.submit_layer(0, _FakeEvent()) ctx.submit_layer(0, _FakeEvent()) self.assertEqual(len(eng.submits), 1) def test_skip_layer_with_no_blocks(self): eng = _FakeEngine() def blocks(L): return None if L == 1 else ([L], [L + 1], [8]) ctx = _ctx(eng, blocks) for L in range(3): ctx.submit_layer(L, None) self.assertEqual(len(eng.submits), 2) # layer 1 (no blocks) skipped def test_finish_waits_all_and_returns_zero(self): eng = _FakeEngine() ctx = _ctx(eng, num_layers=3) for L in range(3): ctx.submit_layer(L, _FakeEvent()) self.assertEqual(ctx.finish(), 0) self.assertEqual(eng.waits, [[101, 102, 103]]) def test_submit_failure_marks_failed_and_stops(self): eng = _FakeEngine(submit_rets=[101, -1, 103]) ctx = _ctx(eng, num_layers=3) for L in range(3): ctx.submit_layer(L, None) self.assertTrue(ctx.failed) self.assertEqual(len(eng.submits), 2) # stopped submitting after the -1 self.assertEqual(ctx.finish(), -1) def test_finish_propagates_wait_failure(self): eng = _FakeEngine(wait_ret=-9) ctx = _ctx(eng, num_layers=1) ctx.submit_layer(0, None) self.assertEqual(ctx.finish(), -9) def test_finish_falls_back_for_notifier_missed_layers(self): # only layer 0 fired by the notifier; finish must SYNCHRONOUSLY transfer the # missing layers 1,2 (e.g. an MTP buffer not hooked) -> all 3 moved, success. eng = _FakeEngine() ctx = _ctx(eng, num_layers=3) ctx.submit_layer(0, None) # notifier fired layer 0 only self.assertEqual(ctx.finish(timeout=0.5), 0) self.assertEqual(len(eng.submits), 3) # layer 0 (overlap) + 1,2 (fallback) self.assertEqual(eng.waits[-1], [101, 102, 103]) def test_failed_submit_still_reported_after_fallback(self): # layer 2 submit fails during the overlap; finish reports -1 even though it # also runs the fallback for any unfired layers. eng = _FakeEngine(submit_rets=[101, -1, 103]) ctx = _ctx(eng, num_layers=3) for L in range(3): ctx.submit_layer(L, None) self.assertEqual(ctx.finish(timeout=0.5), -1) class TestBuildLayerBlocks(unittest.TestCase): def test_addresses_and_lengths(self): from sglang.srt.disaggregation.cp_per_layer_transfer import build_layer_blocks src, dst, lens = build_layer_blocks( 1000, 5000, 64, [[0, 1, 2], [5, 6]], [[10, 11, 12], [20, 21]] ) self.assertEqual(src, [1000 + 0 * 64, 1000 + 5 * 64]) self.assertEqual(dst, [5000 + 10 * 64, 5000 + 20 * 64]) self.assertEqual(lens, [64 * 3, 64 * 2]) # item_len * run length def test_empty(self): from sglang.srt.disaggregation.cp_per_layer_transfer import build_layer_blocks self.assertEqual(build_layer_blocks(1, 2, 8, [], []), ([], [], [])) def test_only_base_ptr_changes_across_layers(self): from sglang.srt.disaggregation.cp_per_layer_transfer import build_layer_blocks s0, d0, l0 = build_layer_blocks(1000, 2000, 64, [[3]], [[7]]) s1, d1, l1 = build_layer_blocks(9000, 8000, 64, [[3]], [[7]]) self.assertEqual(s0, [1000 + 3 * 64]) self.assertEqual(s1, [9000 + 3 * 64]) self.assertEqual(l0, l1) # lengths identical across layers (the invariant) class _MockCtx: def __init__(self): self.submitted = [] self.enqueued = 0 self.finished = False self.failed = False def note_enqueued(self): self.enqueued += 1 def submit_layer(self, layer_id, event): self.submitted.append((layer_id, event)) def finish(self): self.finished = True return 0 def mark_failed(self): self.failed = True class _RecEvent: def __init__(self): self.recorded_stream = "UNSET" self.synced = False def record(self, stream=None): self.recorded_stream = stream def synchronize(self): self.synced = True def _manager(num_workers=0): from sglang.srt.disaggregation.cp_per_layer_transfer import PerLayerTransferManager return PerLayerTransferManager( num_workers=num_workers, event_factory=_RecEvent, current_stream=lambda: "STREAM" ) def _drain(q): import queue as _q items = [] while True: try: items.append(q.get(block=False)) except _q.Empty: break return items class TestPerLayerTransferManager(unittest.TestCase): def test_on_layer_end_enqueues_per_active_ctx_with_recorded_event(self): m = _manager() c1, c2 = _MockCtx(), _MockCtx() m.register("r1", c1) m.register("r2", c2) m.on_layer_end(5) items = _drain(m._q) self.assertEqual(len(items), 2) # one per active context for ctx, layer_id, ev in items: self.assertEqual(layer_id, 5) self.assertEqual(ev.recorded_stream, "STREAM") # event recorded on compute stream def test_worker_step_calls_submit_layer(self): m = _manager() c = _MockCtx() ev = _RecEvent() m._worker_step((c, 3, ev)) self.assertEqual(c.submitted, [(3, ev)]) def test_worker_step_marks_failed_on_exception(self): m = _manager() class _Boom: def __init__(self): self.failed = False def submit_layer(self, layer_id, event): raise RuntimeError("boom") def mark_failed(self): self.failed = True b = _Boom() m._worker_step((b, 0, None)) self.assertTrue(b.failed) def test_finish_pops_and_calls_ctx_finish(self): m = _manager() c = _MockCtx() m.register("r1", c) self.assertTrue(m.has_active()) self.assertEqual(m.finish("r1"), 0) self.assertTrue(c.finished) self.assertFalse(m.has_active()) self.assertEqual(m.finish("r1"), 0) # idempotent after pop def test_on_layer_end_no_active_is_noop(self): m = _manager() m.on_layer_end(0) self.assertEqual(_drain(m._q), []) if __name__ == "__main__": unittest.main()