diff --git a/python/sglang/srt/disaggregation/common/staging_buffer.py b/python/sglang/srt/disaggregation/common/staging_buffer.py index 824a34bbb..baf0d42f4 100644 --- a/python/sglang/srt/disaggregation/common/staging_buffer.py +++ b/python/sglang/srt/disaggregation/common/staging_buffer.py @@ -243,8 +243,12 @@ class StagingAllocator: self.alloc_order.pop(0) if not self.allocations: + # An empty ring makes the entire prior round reusable. Start a + # fresh round at offset zero so the watermark cannot stay stale. + self.round += 1 + self.head = 0 self.watermark_round = self.round - self.watermark_tail = self.head + self.watermark_tail = 0 elif self.alloc_order: off, _, rnd = self.allocations[self.alloc_order[0]] self.watermark_round = rnd diff --git a/python/sglang/srt/disaggregation/common/staging_handler.py b/python/sglang/srt/disaggregation/common/staging_handler.py index e6b77bbca..d9b06eeff 100644 --- a/python/sglang/srt/disaggregation/common/staging_handler.py +++ b/python/sglang/srt/disaggregation/common/staging_handler.py @@ -108,8 +108,17 @@ class DecodeStagingHandler: if receiver is None or not receiver.bootstrap_infos: return key = tuple(str(bi) for bi in receiver.bootstrap_infos) - if key not in self._wm_subscribers: - self._wm_subscribers[key] = (receiver, session_id) + if key in self._wm_subscribers: + return + + self._wm_subscribers[key] = (receiver, session_id) + # Watermark state is per prefill session. Send the current value so a + # new session's first allocation cannot wait on a missed update. + self._send_watermark( + receiver, + session_id, + self.staging_allocator.get_watermark(), + ) def num_writers_for(self, receiver) -> int: """Compute all TP and PP writers expected for a staging chunk.""" @@ -443,26 +452,28 @@ class DecodeStagingHandler: return True def _free_and_send_watermark( - self, alloc_id: int, decode_req: DecodeRequest + self, alloc_id: int, _decode_req: DecodeRequest ) -> None: """Free a staging allocation and broadcast watermark to all prefills.""" self.staging_allocator.free(alloc_id) post_wm = self.staging_allocator.get_watermark() - room = decode_req.req.bootstrap_room - wm_round, wm_tail = post_wm + for receiver, session_id in list(self._wm_subscribers.values()): + self._send_watermark(receiver, session_id, post_wm) + + @staticmethod + def _send_watermark(receiver, session_id: str, watermark) -> None: + """Send one allocator watermark to a registered prefill session.""" + wm_round, wm_tail = watermark wm_round_b = str(wm_round).encode("ascii") wm_tail_b = str(wm_tail).encode("ascii") - for _key, (receiver, session_id) in list(self._wm_subscribers.items()): - sid_b = session_id.encode("ascii") - for bootstrap_info in receiver.bootstrap_infos: - try: - sock, lock = receiver._connect_to_bootstrap_server(bootstrap_info) - with lock: - sock.send_multipart( - [b"WATERMARK", wm_round_b, wm_tail_b, sid_b] - ) - except Exception: - pass + sid_b = session_id.encode("ascii") + for bootstrap_info in receiver.bootstrap_infos: + try: + sock, lock = receiver._connect_to_bootstrap_server(bootstrap_info) + with lock: + sock.send_multipart([b"WATERMARK", wm_round_b, wm_tail_b, sid_b]) + except Exception: + pass def is_watermark_ready( diff --git a/test/registered/unit/disaggregation/test_disaggregation_wire.py b/test/registered/unit/disaggregation/test_disaggregation_wire.py index 1419a5827..7eb936fe0 100644 --- a/test/registered/unit/disaggregation/test_disaggregation_wire.py +++ b/test/registered/unit/disaggregation/test_disaggregation_wire.py @@ -9,7 +9,11 @@ import torch from sglang.srt.disaggregation.base.conn import KVArgs, StateType from sglang.srt.disaggregation.common.conn import CommonKVManager +from sglang.srt.disaggregation.common.staging_buffer import ( + StagingAllocator, +) from sglang.srt.disaggregation.common.staging_handler import ( + DecodeStagingHandler, handle_staging_req, ) from sglang.srt.disaggregation.common.utils import ( @@ -220,6 +224,37 @@ class TestGroupConcurrentContiguous(unittest.TestCase): group_concurrent_contiguous(self._arr([1, 2, 3]), self._arr([1, 2])) +class TestStagingWatermark(unittest.TestCase): + @patch("sglang.srt.disaggregation.common.staging_buffer.StagingBuffer") + def test_empty_ring_restarts_at_zero(self, staging_buffer): + staging_buffer.return_value.data_ptr = 0 + allocator = StagingAllocator(100, "cpu", 0) + alloc_id, _, _ = allocator.assign(60) + + allocator.free(alloc_id) + + self.assertEqual(allocator.get_watermark(), (1, 0)) + self.assertEqual(allocator.assign(70)[1:], (0, 1)) + + def test_new_watermark_subscriber_receives_current_allocator_state(self): + sock = Mock() + bootstrap_info = {"host": "prefill", "port": 7200} + receiver = Mock( + bootstrap_infos=[bootstrap_info], + ) + receiver._connect_to_bootstrap_server.return_value = (sock, threading.Lock()) + handler = object.__new__(DecodeStagingHandler) + handler.staging_allocator = Mock() + handler.staging_allocator.get_watermark.return_value = (3, 0) + handler._wm_subscribers = {} + + handler.register_wm_subscriber(receiver, "session-new") + + sock.send_multipart.assert_called_once_with( + [b"WATERMARK", b"3", b"0", b"session-new"] + ) + + class TestMooncakePPStaging(unittest.TestCase): def test_staging_response_targets_requesting_pp_rank(self): sock = Mock()