diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index 6d806fda3..a3ef89846 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -60,6 +60,7 @@ from sglang.srt.disaggregation.utils import ( get_qsa_pending_state_indices, is_dsv4_c128_online_enabled, is_mla_backend, + is_unadmitted_reject, poll_and_all_reduce, poll_and_all_reduce_pp, poll_and_all_reduce_with_staging, @@ -649,6 +650,13 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin): dispatch happens later, after preallocation and ``send_metadata`` (see ``pop_preallocated``). """ + # See `PrefillBootstrapQueue.add`. A retracted or rebootstrapping + # request owns a host KV backup that `retracted_queue` releases, and by + # this point carries none of the markers `is_unadmitted_reject` reads, + # so take the caller's word for it rather than sniffing. + if not is_retracted and not is_rebootstrap and is_unadmitted_reject(req): + self.scheduler.retire_unadmitted_request(req) + return if self._check_if_req_exceed_kv_capacity(req): return diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index 56604a281..837b0c014 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -52,6 +52,7 @@ from sglang.srt.disaggregation.utils import ( is_aborted, is_dsv4_c128_online_enabled, is_mla_backend, + is_unadmitted_reject, poll_and_all_reduce_attn_cp_tp_group, poll_and_all_reduce_pp, prepare_abort, @@ -397,6 +398,13 @@ class PrefillBootstrapQueue: return True def add(self, req: Req, num_kv_heads: int) -> None: + # Rejected at intake: `set_finish_with_abort` left the verdict in + # `to_finish`, which `finished()` does not read, and swapped the prompt + # for a one-token stub. Bootstrapping it costs a handshake, a metadata + # buffer and a forward pass before anything unwinds it. + if is_unadmitted_reject(req): + self.scheduler.retire_unadmitted_request(req) + return if not self.create_sender(req, num_kv_heads): return self.queue.append(req) diff --git a/python/sglang/srt/disaggregation/utils.py b/python/sglang/srt/disaggregation/utils.py index c6bd1080c..3498b99e6 100644 --- a/python/sglang/srt/disaggregation/utils.py +++ b/python/sglang/srt/disaggregation/utils.py @@ -1704,6 +1704,35 @@ def prepare_abort(req: Req, error_message: str, status_code=None): req.logprob.input_token_ids_logprobs_idx = [] +def is_unadmitted_reject(req: Req) -> bool: + """A request rejected at intake, before it acquired anything. + + A preempted or resumed request can also carry a pending abort -- "Abort + method 3" marks a *running* request and `filter_batch` does not drop it, + since `finished()` is still False -- and its queue owns the release of + whatever it still holds. + + `req.is_retracted` catches the two re-entries that declare nothing: + priority preemption and the pause/retract-all path both requeue through a + bare `_add_request_to_queue`. `release_req` always calls + `reset_for_retract`, which sets it, and its clear sites all run downstream + of these doors. The resource markers stay as a second line of defence -- + on their own they miss a `seqlen <= 1` preemption, whose KV is already + freed and whose `retraction_backup` was never taken. + + `DecodePreallocQueue.add` still gates on its own `is_retracted` / + `is_rebootstrap` parameters as well, since they state the caller's intent + rather than inferring it. + """ + return is_aborted(req) and not ( + req.is_retracted + or req.kv.holds_kv + or req.kv.holds_mamba + or req.metadata_buffer_index >= 0 + or req.kv.retraction_backup is not None + ) + + def is_aborted(req: Req) -> bool: from sglang.srt.managers.schedule_batch import FINISH_ABORT diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index eaee6df66..d63f4e8d6 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -3198,6 +3198,30 @@ class Scheduler( self._retry_storage_prefetch(req) return True + def retire_unadmitted_request(self, req: Req) -> None: + """Finish a request the disaggregation queues rejected at their door.""" + # `create_req` marks a streaming session in-flight, and the pre-abort + # detach lives in `StreamingSession.find_active_slot`, which only runs + # while scheduling; a session left in-flight rejects every later request. + if req.session is not None and req.session.streaming: + req.session.abort_req() + req.session = None + # `beam_coordinator.validate_and_init` counts the group in ahead of the + # checks that reject; no-op when the request has no group. + self.beam_coordinator.retire_group(req) + # PREFILL runs `_prefetch_kvcache` before its door, so even the + # one-token stub is registered with the cache by now: + # `prefetch_from_storage` arms the paced-retry set for this attempt's + # cache handle. Only a `finish`/ABORT and a `waiting_queue` sweep clear + # that, and a retired request reaches neither. + self._release_aborted_request(req) + # `update_finish_state` returns early once `finished()`, so an already + # set `finished_reason` is what the client receives; report the same. + reason = req.finished_reason or req.to_finish + req.time_stats.trace_ctx.abort(abort_info={"reason": reason.message}) + req.update_finish_state() + self.output_streamer.stream_output([req], req.return_logprob) + def _add_request_to_queue(self, req: Req, is_retracted: bool = False): if not self._set_or_validate_priority(req): return diff --git a/test/registered/unit/disaggregation/test_admission_abort_not_enqueued.py b/test/registered/unit/disaggregation/test_admission_abort_not_enqueued.py new file mode 100644 index 000000000..f5a8b23b5 --- /dev/null +++ b/test/registered/unit/disaggregation/test_admission_abort_not_enqueued.py @@ -0,0 +1,326 @@ +"""A request rejected at intake must not be admitted to a PD handoff. + +`Req.set_finish_with_abort()` records the rejection in `to_finish` (not +`finished_reason`) and replaces the prompt with a one-token stub, so +`req.finished()` stays False and the disaggregation admission queues let it +through: it completes a bootstrap handshake, reserves a metadata buffer, +initialises an RDMA sender and runs a forward pass on the stub before anything +unwinds it, and the decode worker sits on it until the transfer timeout. + +Only disaggregation is affected. In NULL mode the stub costs one cheap forward +pass and the batch boundary returns the 400, so that path is deliberately left +alone -- `test_null_mode_is_deliberately_untouched` pins that. + +The door must also not retire a *re-entering* request. "Abort method 3" sets +`to_finish` on a running request and `filter_batch` does not drop it, so a +retracted, preempted or resumed request can arrive carrying one while its queue +still owes a release. The prefill door detects that from the resources the +request still holds; the decode door cannot -- a retracted request holds none of +them by then -- so it gates on its own `is_retracted` / `is_rebootstrap` flags. +""" + +import unittest +from array import array +from http import HTTPStatus +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +from sglang.srt.disaggregation.decode import DecodePreallocQueue +from sglang.srt.disaggregation.prefill import PrefillBootstrapQueue +from sglang.srt.disaggregation.utils import DisaggregationMode, is_unadmitted_reject +from sglang.srt.managers.schedule_batch import FINISH_ABORT, Req +from sglang.srt.managers.scheduler import Scheduler +from sglang.srt.runtime_context import publish, reset_context +from sglang.srt.sampling.sampling_params import SamplingParams +from sglang.srt.server_args import ServerArgs +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=3, suite="base-a-test-cpu") + +ERROR_MSG = ( + "Input length (1500 tokens) exceeds the maximum allowed length (1018 tokens)." +) + + +def _make_req(prompt_len=1500, session=None): + req = Req( + rid="admission-abort", + origin_input_text="", + origin_input_ids=array("q", list(range(prompt_len))), + sampling_params=SamplingParams(max_new_tokens=8), + session=session, + ) + req.time_stats.trace_ctx = MagicMock() + return req + + +def _make_scheduler(): + """One stub for both halves. + + `retire_unadmitted_request` is a spy wrapping the real method, so a door + test still asserts the call *and* runs the body -- otherwise the doors and + the retirement are only ever tested apart, and a step missing from the + retirement passes every door test. + """ + sched = SimpleNamespace( + output_streamer=MagicMock(), + beam_coordinator=MagicMock(), + _release_aborted_request=MagicMock(), + ) + sched.retire_unadmitted_request = MagicMock( + side_effect=lambda req: Scheduler.retire_unadmitted_request(sched, req) + ) + return sched + + +def _prefill_queue(sched): + q = SimpleNamespace( + scheduler=sched, + queue=[], + create_sender=MagicMock(return_value=True), + ) + return q + + +def _decode_queue(sched): + q = SimpleNamespace( + scheduler=sched, + retracted_queue=[], + pending_reqs=[], + _check_if_req_exceed_kv_capacity=MagicMock(return_value=False), + _create_receiver_and_enqueue=MagicMock( + return_value=SimpleNamespace(kv_receiver=MagicMock()) + ), + _resolve_prefill_dp_rank=MagicMock(return_value=0), + ) + return q + + +def _admit_prefill(q, req): + PrefillBootstrapQueue.add(q, req, 8) + + +def _admit_decode(q, req, **kw): + with patch( + "sglang.srt.disaggregation.decode._is_fake_transfer", return_value=False + ): + DecodePreallocQueue.add(q, req, **kw) + + +class TestAdmissionAbortNotEnqueued(CustomTestCase): + def setUp(self): + reset_context() + self.addCleanup(reset_context) + publish(ServerArgs(model_path="dummy"), role="tokenizer") + # set_finish_with_abort() logs on TP rank 0, and ParallelContext.tp_rank + # reads through to the live process group -- none exists here. + p = patch( + "sglang.srt.managers.schedule_batch.get_parallel", + return_value=SimpleNamespace(tp_rank=0), + ) + p.start() + self.addCleanup(p.stop) + + def test_rejected_request_is_retired_at_each_pd_door(self): + for door in ("prefill", "decode"): + with self.subTest(door=door): + req = _make_req() + req.set_finish_with_abort(ERROR_MSG) + # Precondition: the rejection is pending, so `finished()` -- + # what every admission gate keys off -- is still False. + self.assertIsInstance(req.to_finish, FINISH_ABORT) + self.assertFalse(req.finished()) + self.assertEqual(len(req.origin_input_ids), 1) + + sched = _make_scheduler() + if door == "prefill": + q = _prefill_queue(sched) + _admit_prefill(q, req) + q.create_sender.assert_not_called() + self.assertEqual(q.queue, []) + else: + q = _decode_queue(sched) + _admit_decode(q, req) + q._create_receiver_and_enqueue.assert_not_called() + self.assertEqual(q.pending_reqs, []) + sched.retire_unadmitted_request.assert_called_once_with(req) + + def test_reentering_prefill_request_is_not_retired(self): + """A preempted prefill requeue must reach its queue, which owns release. + + Preemption runs `release_req`, so KV is gone by then; what survives is + the metadata buffer (`finalize_bootstrap` allocated it before the + request ever entered the running batch) and the host retraction backup. + """ + for marker in ("metadata_buffer_index", "retraction_backup"): + with self.subTest(marker=marker): + req = _make_req(prompt_len=16) + req.to_finish = FINISH_ABORT( + "Aborted by AbortReq.", HTTPStatus.SERVICE_UNAVAILABLE + ) + # release_req already freed the KV row. + self.assertFalse(req.kv.holds_kv) + if marker == "metadata_buffer_index": + req.metadata_buffer_index = 7 + else: + req.kv.retraction_backup = object() + self.assertFalse(is_unadmitted_reject(req)) + + sched = _make_scheduler() + q = _prefill_queue(sched) + _admit_prefill(q, req) + self.assertEqual(q.queue, [req]) + sched.retire_unadmitted_request.assert_not_called() + self.assertIsInstance(req.to_finish, FINISH_ABORT) + + def test_retracted_decode_request_is_not_retired(self): + """The decode door cannot sniff this one, so it must trust its flags. + + By the time `retract_decode` requeues, the request carries none of the + markers `is_unadmitted_reject` reads: `release_req` nulled + `kv.req_pool_idx`, `reset_for_retract` nulled `kv.mamba_pool_idx`, the + decode metadata buffer lives on `DecodeRequest` rather than `Req`, and + `add()` itself clears `retraction_mb_id`. Retiring it here would strand + the host pages `release_req` allocated, which only `retraction_restore` + or `retraction_discard` free. + + `backup=None` is the short-sequence case: `retraction_backup()` returns + early for `seqlen <= 1` without setting it, so the resource predicate + alone says "unadmitted" and only the flag keeps the request safe. + """ + for flag in ("is_retracted", "is_rebootstrap"): + for backup in (object(), None): + with self.subTest(flag=flag, has_backup=backup is not None): + req = _make_req(prompt_len=16) + req.to_finish = FINISH_ABORT( + "Aborted by AbortReq.", HTTPStatus.SERVICE_UNAVAILABLE + ) + # Exactly what retract_decode leaves behind. + self.assertFalse(req.kv.holds_kv) + self.assertFalse(req.kv.holds_mamba) + self.assertEqual(req.metadata_buffer_index, -1) + req.kv.retraction_backup = backup + + sched = _make_scheduler() + q = _decode_queue(sched) + _admit_decode(q, req, **{flag: True}) + sched.retire_unadmitted_request.assert_not_called() + self.assertIsInstance(req.to_finish, FINISH_ABORT) + if flag == "is_retracted": + self.assertEqual(q.retracted_queue, [req]) + else: + q._create_receiver_and_enqueue.assert_called_once() + + def test_preempted_reentry_is_not_retired(self): + """Preemption requeues through a bare `_add_request_to_queue`. + + The door sees no flag at all, and `release_req` -> `reset_for_retract` + is the only thing that marks the request. The resource markers miss + this shape on their own: the KV row is already freed and + `retraction_backup` is never taken for `seqlen <= 1`. + """ + for door in ("prefill", "decode"): + with self.subTest(door=door): + req = _make_req(prompt_len=16) + req.to_finish = FINISH_ABORT( + "Aborted by AbortReq.", HTTPStatus.SERVICE_UNAVAILABLE + ) + req.is_retracted = True + # Exactly what a short preemption leaves behind. + self.assertFalse(req.kv.holds_kv) + self.assertIsNone(req.kv.retraction_backup) + self.assertEqual(req.metadata_buffer_index, -1) + self.assertFalse(is_unadmitted_reject(req)) + + sched = _make_scheduler() + if door == "prefill": + q = _prefill_queue(sched) + _admit_prefill(q, req) + self.assertEqual(q.queue, [req]) + else: + q = _decode_queue(sched) + _admit_decode(q, req) + q._create_receiver_and_enqueue.assert_called_once() + sched.retire_unadmitted_request.assert_not_called() + + def test_valid_request_still_admitted(self): + for door in ("prefill", "decode"): + with self.subTest(door=door): + req = _make_req(prompt_len=16) + sched = _make_scheduler() + if door == "prefill": + q = _prefill_queue(sched) + _admit_prefill(q, req) + self.assertEqual(q.queue, [req]) + else: + q = _decode_queue(sched) + _admit_decode(q, req) + q._create_receiver_and_enqueue.assert_called_once() + sched.retire_unadmitted_request.assert_not_called() + + def test_null_mode_is_deliberately_untouched(self): + """The bug is disaggregation-only; NULL mode must keep enqueuing. + + There the one-token stub costs a cheap forward pass and the batch + boundary promotes `to_finish` into the 400. Adding a guard here would + skip `StreamingSession.find_active_slot`'s pre-abort detach, which only + runs while scheduling. + """ + req = _make_req() + req.set_finish_with_abort(ERROR_MSG) + sched = SimpleNamespace( + disaggregation_mode=DisaggregationMode.NULL, + waiting_queue=[], + processed_tokens_counter=0, + _set_or_validate_priority=MagicMock(return_value=True), + _abort_on_queued_limit=MagicMock(return_value=False), + _prefetch_kvcache=MagicMock(), + ) + Scheduler._add_request_to_queue(sched, req) + self.assertEqual(sched.waiting_queue, [req]) + self.assertIsInstance(req.to_finish, FINISH_ABORT) + + def test_retire_detaches_a_streaming_session(self): + """Otherwise the session stays in-flight forever. + + `create_req` marks it in-flight and the pre-abort detach lives in + `StreamingSession.find_active_slot`, which a retired request never + reaches; a stuck flag fails every later request on that session. + """ + session = MagicMock() + session.streaming = True + req = _make_req(session=session) + req.set_finish_with_abort(ERROR_MSG) + + sched = _make_scheduler() + Scheduler.retire_unadmitted_request(sched, req) + + # PREFILL arms the cache's paced-retry set via `_prefetch_kvcache` + # before the door, and only this call clears it again. + sched._release_aborted_request.assert_called_once_with(req) + session.abort_req.assert_called_once() + self.assertIsNone(req.session) + sched.beam_coordinator.retire_group.assert_called_once_with(req) + req.time_stats.trace_ctx.abort.assert_called_once() + # The original 400 survives, rather than being replaced downstream. + self.assertIsInstance(req.finished_reason, FINISH_ABORT) + self.assertEqual(req.finished_reason.status_code, HTTPStatus.BAD_REQUEST) + self.assertIsNone(req.to_finish) + sched.output_streamer.stream_output.assert_called_once() + + def test_retire_leaves_a_non_streaming_session_alone(self): + session = MagicMock() + session.streaming = False + req = _make_req(session=session) + req.set_finish_with_abort(ERROR_MSG) + + sched = _make_scheduler() + Scheduler.retire_unadmitted_request(sched, req) + + session.abort_req.assert_not_called() + self.assertIs(req.session, session) + + +if __name__ == "__main__": + unittest.main()