diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index d0293fba3..4be51ca1f 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -725,7 +725,9 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin): error_msg = f"Could not fetch prefill parallel info from {bootstrap_addr} after {count} attempts" logger.error(error_msg) for decode_req in reqs: - decode_req.kv_receiver.abort() + # kv_receiver may be None from a prior self.queue cleanup + if decode_req.kv_receiver is not None: + decode_req.kv_receiver.abort() del self._ensure_retry_count[bootstrap_addr] del self._ensure_last_attempt_time[bootstrap_addr] else: @@ -837,6 +839,14 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin): failed_reqs.append(decode_req) indices_to_remove.add(i) + # DecodeRequest is shared between self.queue and self.pending_reqs; + # drop failed reqs from both + if failed_reqs: + failed_ids = {id(r) for r in failed_reqs} + self.pending_reqs = [ + r for r in self.pending_reqs if id(r) not in failed_ids + ] + # HiSparse physical constraint: max requests by device buffer capacity. # Each admitted req needs padded_buffer_size from hisparse device pool. # waiting_queue reqs already have device buffers (allocated in admit_request_direct), diff --git a/test/registered/unit/disaggregation/test_decode_queue_cleanup.py b/test/registered/unit/disaggregation/test_decode_queue_cleanup.py index 03184646d..6cb9f53da 100644 --- a/test/registered/unit/disaggregation/test_decode_queue_cleanup.py +++ b/test/registered/unit/disaggregation/test_decode_queue_cleanup.py @@ -66,6 +66,70 @@ class TestDecodeQueueCleanup(CustomTestCase): [req], req.return_logprob ) + def test_prealloc_abort_also_drops_from_pending_reqs(self): + # Same DecodeRequest lives in both queue and pending_reqs (add() slow + # path). Aborting must drop it from both, and compare by identity since + # DecodeRequest's dataclass __eq__ would compare the tensor receiver. + class BadEqReceiver(FakeReceiver): + def __eq__(self, other): + raise TypeError("use identity comparison, not value equality") + + __hash__ = object.__hash__ + + receiver = BadEqReceiver() + req = SimpleNamespace( + rid="abort-shared", + finished_reason=FINISH_ABORT("aborted"), + return_logprob=False, + ) + decode_req = SimpleNamespace(req=req, kv_receiver=receiver) + + queue = DecodePreallocQueue.__new__(DecodePreallocQueue) + queue.queue = [decode_req] + queue.pending_reqs = [decode_req] # same object, dual ownership + 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) + + scheduler = MagicMock() + scheduler.running_batch.reqs = [] + scheduler.enable_priority_scheduling = False + scheduler.enable_hisparse = False + scheduler.output_streamer = MagicMock() + queue.scheduler = scheduler + + # Must not raise on the receiver __eq__ above. + preallocated, failed = queue.pop_preallocated() + + self.assertEqual(preallocated, []) + self.assertEqual(failed, [decode_req]) + self.assertEqual(queue.queue, []) + self.assertTrue(all(r is not decode_req for r in queue.pending_reqs)) + self.assertIsNone(decode_req.kv_receiver) + + def test_ensure_prefill_info_tolerates_cleared_receiver(self): + # A req whose kv_receiver was already cleared must not crash on .abort(). + queue = DecodePreallocQueue.__new__(DecodePreallocQueue) + queue._max_ensure_retries = 1 + queue._ensure_retry_interval = 0 + queue._ensure_retry_count = {"127.0.0.1:11500": 0} + queue._ensure_last_attempt_time = {} + queue.kv_manager = MagicMock() + queue.kv_manager.try_ensure_parallel_info.return_value = False + + cleared_req = SimpleNamespace( + req=SimpleNamespace(rid="cleared"), kv_receiver=None + ) + addr_to_reqs = {"127.0.0.1:11500": [cleared_req]} + + ready, remaining = queue._ensure_prefill_info(addr_to_reqs) + + self.assertEqual(ready, {}) + self.assertEqual(remaining, []) + @patch("sglang.srt.disaggregation.decode.release_kv_cache") @patch("sglang.srt.disaggregation.decode.prepare_abort") @patch("sglang.srt.disaggregation.decode.poll_and_all_reduce")