From 18989f3d4812f144d862a63c873bf1f93034ebc0 Mon Sep 17 00:00:00 2001 From: luoroger37 Date: Fri, 12 Jun 2026 18:20:24 +0800 Subject: [PATCH] [PD] Fix resource leak on prealloc/transfer abort and idle check (#28022) --- python/sglang/srt/disaggregation/decode.py | 4 + python/sglang/srt/managers/scheduler.py | 1 + .../test_decode_queue_cleanup.py | 151 ++++++++++++++++++ 3 files changed, 156 insertions(+) create mode 100644 test/registered/unit/disaggregation/test_decode_queue_cleanup.py diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index 0af05cd01..1ecd04ae1 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -815,6 +815,8 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin): [decode_req.req], decode_req.req.return_logprob, ) + decode_req.kv_receiver.clear() + decode_req.kv_receiver = None failed_reqs.append(decode_req) indices_to_remove.add(i) @@ -1651,6 +1653,8 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin): self.scheduler.hisparse_coordinator.request_finished(decode_req.req) # release pre-allocated kv cache, but don't insert into the tree since it's failed release_kv_cache(decode_req.req, self.tree_cache, is_insert=False) + decode_req.kv_receiver.clear() + decode_req.kv_receiver = None indices_to_remove.add(i) if self.scheduler.metrics_reporter.enable_metrics: self.scheduler.metrics_collector.increment_transfer_failed_reqs() diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 18e6a5932..99fdf2c28 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -3380,6 +3380,7 @@ class Scheduler( if self.disaggregation_mode == DisaggregationMode.DECODE: idle &= len(self.disagg_decode_prealloc_queue.queue) == 0 + idle &= len(self.disagg_decode_prealloc_queue.retracted_queue) == 0 idle &= len(self.disagg_decode_transfer_queue.queue) == 0 if self.decode_offload_manager is not None: idle &= len(self.decode_offload_manager.ongoing_offload) == 0 diff --git a/test/registered/unit/disaggregation/test_decode_queue_cleanup.py b/test/registered/unit/disaggregation/test_decode_queue_cleanup.py new file mode 100644 index 000000000..03184646d --- /dev/null +++ b/test/registered/unit/disaggregation/test_decode_queue_cleanup.py @@ -0,0 +1,151 @@ +import unittest +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +from sglang.srt.disaggregation.base import KVPoll +from sglang.srt.disaggregation.decode import ( + DecodePreallocQueue, + DecodeTransferQueue, + HiCacheRestoreResult, +) +from sglang.srt.disaggregation.utils import DisaggregationMode +from sglang.srt.managers.schedule_batch import FINISH_ABORT +from sglang.srt.managers.scheduler import Scheduler +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + + +class FakeReceiver: + def __init__(self): + self.clear_called = False + + def clear(self): + self.clear_called = True + + def failure_exception(self): + return None + + +class TestDecodeQueueCleanup(CustomTestCase): + def test_prealloc_abort_clears_receiver_before_removing_request(self): + receiver = FakeReceiver() + req = SimpleNamespace( + rid="abort-prealloc", + 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 = [] + 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 + + preallocated, failed = queue.pop_preallocated() + + self.assertEqual(preallocated, []) + self.assertEqual(failed, [decode_req]) + self.assertEqual(queue.queue, []) + self.assertTrue(receiver.clear_called) + self.assertIsNone(decode_req.kv_receiver) + scheduler.output_streamer.stream_output.assert_called_once_with( + [req], req.return_logprob + ) + + @patch("sglang.srt.disaggregation.decode.release_kv_cache") + @patch("sglang.srt.disaggregation.decode.prepare_abort") + @patch("sglang.srt.disaggregation.decode.poll_and_all_reduce") + def test_transfer_failure_clears_receiver_before_removing_request( + self, mock_poll, mock_prepare_abort, mock_release_kv_cache + ): + receiver = FakeReceiver() + req = SimpleNamespace( + rid="failed-transfer", + bootstrap_room=7, + return_logprob=False, + ) + decode_req = SimpleNamespace( + req=req, + kv_receiver=receiver, + metadata_buffer_index=3, + hicache_restore_status=HiCacheRestoreResult.READY, + ) + + queue = DecodeTransferQueue.__new__(DecodeTransferQueue) + queue.queue = [decode_req] + queue.enable_staging = False + queue.gloo_group = MagicMock() + queue.req_to_metadata_buffer_idx_allocator = MagicMock() + queue.tp_rank = 0 + queue.tree_cache = MagicMock() + queue.metadata_buffers = SimpleNamespace(bootstrap_room=[None] * 4) + queue.spec_algorithm = MagicMock() + queue.spec_algorithm.is_none.return_value = True + queue._clean_hicache_prefetch_resources = MagicMock() + + scheduler = MagicMock() + scheduler.enable_decode_hicache = False + scheduler.enable_hisparse = False + scheduler.output_streamer = MagicMock() + scheduler.metrics_reporter.enable_metrics = False + queue.scheduler = scheduler + + mock_poll.return_value = [KVPoll.Failed] + + transferred = queue.pop_transferred() + + self.assertEqual(transferred, []) + self.assertEqual(queue.queue, []) + self.assertTrue(receiver.clear_called) + self.assertIsNone(decode_req.kv_receiver) + queue.req_to_metadata_buffer_idx_allocator.free.assert_called_once_with(3) + scheduler.output_streamer.stream_output.assert_called_once_with( + [req], req.return_logprob + ) + mock_prepare_abort.assert_called_once() + mock_release_kv_cache.assert_called_once_with( + req, queue.tree_cache, is_insert=False + ) + + def test_retracted_decode_requests_keep_scheduler_non_idle(self): + scheduler = Scheduler.__new__(Scheduler) + scheduler.running_batch = MagicMock() + scheduler.running_batch.is_empty.return_value = True + scheduler.chunked_req = None + scheduler.dllm_manager = MagicMock() + scheduler.dllm_manager.any_staging_reqs.return_value = False + scheduler.last_batch = None + scheduler.cur_batch = None + scheduler.enable_overlap = False + scheduler.ps = SimpleNamespace(pp_size=1) + scheduler.running_mbs = [] + scheduler.waiting_queue = [] + scheduler.grammar_manager = SimpleNamespace(grammar_queue=[]) + scheduler.disaggregation_mode = DisaggregationMode.DECODE + scheduler.disagg_decode_prealloc_queue = SimpleNamespace( + queue=[], retracted_queue=[object()] + ) + scheduler.disagg_decode_transfer_queue = SimpleNamespace(queue=[]) + scheduler.decode_offload_manager = None + scheduler.enable_hisparse = False + scheduler.enable_hierarchical_cache = False + + self.assertFalse(scheduler.is_fully_idle()) + + +if __name__ == "__main__": + unittest.main()