diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index 5de5ba984..4cda16dcd 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -869,17 +869,8 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin): if not self.queue: return - # Still poll if any receiver was aborted, otherwise it stays stuck. - if ( - self.pp_size <= 1 - and all(decode_req.waiting_for_input for decode_req in self.queue) - and not any( - decode_req.kv_receiver.conclude_state == KVPoll.Failed - for decode_req in self.queue - ) - ): - return - + # Receiver polling observes asynchronous failures while KV allocation + # is blocked. if self.pp_size > 1: polls = poll_and_all_reduce_pp( (decode_req.req.rid for decode_req in self.queue), diff --git a/test/registered/unit/disaggregation/test_decode_queue_cleanup.py b/test/registered/unit/disaggregation/test_decode_queue_cleanup.py index 7b5b9f8a0..7c979edce 100644 --- a/test/registered/unit/disaggregation/test_decode_queue_cleanup.py +++ b/test/registered/unit/disaggregation/test_decode_queue_cleanup.py @@ -23,6 +23,7 @@ register_cpu_ci(est_time=5, suite="base-a-test-cpu") class FakeReceiver: def __init__(self): self.clear_called = False + self.conclude_state = None def clear(self): self.clear_called = True @@ -94,18 +95,22 @@ class TestDecodeQueueCleanup(CustomTestCase): receiver = FakeReceiver() req = SimpleNamespace( rid="abort-prealloc", - finished_reason=FINISH_ABORT("aborted"), + bootstrap_room=42, + finished_reason=None, return_logprob=False, ) - decode_req = SimpleNamespace(req=req, kv_receiver=receiver) + decode_req = SimpleNamespace( + req=req, kv_receiver=receiver, waiting_for_input=True + ) queue = DecodePreallocQueue.__new__(DecodePreallocQueue) queue.pp_size = 1 + queue.tp_rank = 0 + queue.gloo_group = object() queue.queue = [decode_req] queue.pending_reqs = [] queue.retracted_queue = [] queue._resolve_pending_reqs = MagicMock() - queue._update_handshake_waiters = MagicMock() queue._uses_swa_tail_prealloc = MagicMock(return_value=False) queue._allocatable_token_budgets = MagicMock(return_value=0) queue._hicache_pending_restore_tokens = MagicMock(return_value=0) @@ -114,16 +119,23 @@ class TestDecodeQueueCleanup(CustomTestCase): scheduler.running_batch.reqs = [] scheduler.enable_priority_scheduling = False scheduler.enable_hisparse = False + scheduler.metrics_reporter.enable_metrics = False scheduler.output_streamer = MagicMock() queue.scheduler = scheduler - preallocated, failed = queue.pop_preallocated() + with patch( + "sglang.srt.disaggregation.decode.poll_and_all_reduce", + return_value=[KVPoll.Failed], + ) as poll: + preallocated, failed = queue.pop_preallocated() + poll.assert_called_once_with([receiver], queue.gloo_group) self.assertEqual(preallocated, []) self.assertEqual(failed, [decode_req]) self.assertEqual(queue.queue, []) self.assertTrue(receiver.clear_called) self.assertIsNone(decode_req.kv_receiver) + self.assertIsInstance(req.finished_reason, FINISH_ABORT) scheduler.output_streamer.stream_output.assert_called_once_with( [req], req.return_logprob )