diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index e6dfaf99f..7a5342d28 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -3348,6 +3348,11 @@ class Scheduler( # it. Drop the marker once the request is actually gone. if req.finished() or not req.kv.holds_kv: self._pending_chunked_abort_req = None + return + # The request moved to another scheduler queue after abort_request + # deferred it, so retry against its current location. + self._pending_chunked_abort_req = None + self.abort_request(AbortReq(rid=req.rid)) return prepare_abort(req, "Aborted") diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index c5a5e24cc..cc7d03d4b 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -233,6 +233,9 @@ class ReqState: last_completion_tokens: int = 1 ttft_observed: bool = False + dispatched: bool = False + abort_sent: bool = False + # For streaming output last_output_offset: int = 0 @@ -799,6 +802,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): ) self._init_req_state(obj, request) + request_rids = {obj.rid} if obj.is_single else set(obj.rid) try: if get_disagg().language_only: self._handle_epd_disaggregation_encode_request(obj) @@ -822,17 +826,19 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): async for response in self._wait_one_response(obj, request): yield response else: - async for response in self._handle_batch_request(obj, request): + async for response in self._handle_batch_request( + obj, request, request_rids + ): yield response except BaseException: # _init_req_state created a rid_to_state entry per (sub-)request up # front. The normal remover is the scheduler-response path # (_handle_batch_output), so a failure *before* a request reaches the # scheduler -- e.g. input-length validation rejecting an over-context - # request -- would otherwise leak those entries forever. Drop any that - # are still pending; entries already removed on the normal completion - # path are left untouched (pop is a no-op). - self._discard_pending_req_states(obj) + # request -- would otherwise leak those entries forever. Drop + # undelivered states, but abort dispatched requests for scheduler-side + # cleanup. + self._release_req_states_on_failure(request_rids) raise def _detect_input_format( @@ -1571,6 +1577,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): time_stats = tokenized_obj.time_stats tokenized_obj.wrap_pickle_fields() self._dispatch_to_scheduler(tokenized_obj) + self._mark_state_dispatched(tokenized_obj.rid) dispatched = True tokenized_obj.time_stats = time_stats tokenized_obj.time_stats.set_api_server_dispatch_finish_time() @@ -1578,6 +1585,16 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): if not dispatched: self.cuda_vmm_feature_transport.cancel_for_dispatch(prepared_mm_items) + def _mark_state_dispatched(self, rid: str): + """Record that *rid* reached the scheduler. + + Only dispatched requests are aborted (not discarded) by the + handler-failure cleanup; see _release_req_states_on_failure. + """ + state = self.rid_to_state.get(rid) + if state is not None: + state.dispatched = True + async def _send_batch_request( self, tokenized_objs: List[ @@ -1605,6 +1622,8 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): batch_req = BatchTokenizedEmbeddingReqInput(batch=tokenized_objs) self._dispatch_to_scheduler(batch_req) + for tokenized_obj in tokenized_objs: + self._mark_state_dispatched(tokenized_obj.rid) dispatched = True for tokenized_obj, time_stat in zip(tokenized_objs, time_stats): tokenized_obj.time_stats = time_stat @@ -1819,7 +1838,10 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): self, obj: Union[GenerateReqInput, EmbeddingReqInput], request: Optional[fastapi.Request] = None, + request_rids: Optional[set[str]] = None, ): + if request_rids is None: + request_rids = set(obj.rid) batch_size = obj.batch_size generators = [] @@ -1885,6 +1907,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): tokenized_obj.sampling_params.max_new_tokens = 0 tokenized_obj.stream = False self._init_req_state(tmp_obj) + request_rids.add(tmp_obj.rid) await self._send_one_request(tokenized_obj) await self._wait_one_response(tmp_obj, request).__anext__() @@ -1901,6 +1924,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): ] tokenized_obj.rid = tmp_obj.regenerate_rid() self._init_req_state(tmp_obj) + request_rids.add(tmp_obj.rid) state = self.rid_to_state[tmp_obj.rid] tokenized_obj.time_stats = state.time_stats if tmp_obj.return_prompt_token_ids: @@ -1970,14 +1994,21 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): if not abort_all and not rid: logger.warning("Ignore abort_request with empty rid and abort_all=False") return - if ( - not abort_all - and get_serving().tokenizer_worker_num == 1 - and rid not in self.rid_to_state - ): - return + state = None if abort_all else self.rid_to_state.get(rid) + if not abort_all: + if state is not None: + if state.abort_sent: + return + state.abort_sent = True + elif get_serving().tokenizer_worker_num == 1: + return req = AbortReq(rid=rid, abort_all=abort_all) - self._dispatch_to_scheduler(req) + try: + self._dispatch_to_scheduler(req) + except BaseException: + if state is not None: + state.abort_sent = False + raise if self.enable_metrics: # TODO: also use custom_labels from the request self.metrics_collector.observe_one_aborted_request( @@ -2140,10 +2171,9 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): # Abort the request if the client is disconnected. async def abort_request(): await asyncio.sleep(2) - if obj.is_single: - self.abort_request(obj.rid) - else: - for rid in obj.rid: + rids = [obj.rid] if obj.is_single else obj.rid + for rid in rids: + if rid in self.rid_to_state: self.abort_request(rid) background_tasks = BackgroundTasks() @@ -3455,19 +3485,23 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): time_stats.init_trace_ctx(rid, bootstrap_room, external_trace_header) time_stats.set_created_time(created_time) - def _discard_pending_req_states(self, obj): - """Drop rid_to_state entries created by _init_req_state for *obj*. + def _release_req_states_on_failure(self, rids: Iterable[str]): + """Release rid_to_state entries created for a failed handler. - Safe to call after a partial/failed dispatch: only entries still present - are removed, and the scheduler-response path looks up state with - ``.get(...)`` so a later output for a discarded rid is ignored, not fatal. + Undelivered states are removed locally. Dispatched requests are aborted + and retained until the scheduler response removes them. """ - if not hasattr(obj, "is_single") or obj.is_single: - rids = [obj.rid] - else: - rids = obj.rid for rid in rids: - self.rid_to_state.pop(rid, None) + state = self.rid_to_state.get(rid) + if state is None: + continue + if state.dispatched: + try: + self.abort_request(rid) + except Exception: + logger.exception("Failed to abort request %s during cleanup", rid) + else: + del self.rid_to_state[rid] def _should_dispatch_to_encoder( self, obj: Union[GenerateReqInput, EmbeddingReqInput] diff --git a/test/registered/unit/managers/test_scheduler_chunked_abort_race.py b/test/registered/unit/managers/test_scheduler_chunked_abort_race.py new file mode 100644 index 000000000..38f514383 --- /dev/null +++ b/test/registered/unit/managers/test_scheduler_chunked_abort_race.py @@ -0,0 +1,69 @@ +"""Tests for deferred chunked-prefill aborts.""" + +import unittest +from types import SimpleNamespace +from unittest.mock import Mock + +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase, maybe_stub_sgl_kernel + +maybe_stub_sgl_kernel() + +from sglang.srt.managers.scheduler import Scheduler # noqa: E402 + +register_cpu_ci(est_time=10, suite="base-a-test-cpu") + + +class _FakeReq: + """Minimal stand-in for Req: only the fields the abort paths touch.""" + + def __init__(self, rid: str): + self.rid = rid + # Mirrors Req.kv; the abort paths read only these two predicates. + self.kv = SimpleNamespace(holds_kv=True, holds_mamba=False) + self.to_finish = None + self._finished = False + + def finished(self): + return self._finished + + +def _make_scheduler(pending_req, *, chunked_req, running_reqs) -> Scheduler: + sched = Scheduler.__new__(Scheduler) + sched.chunked_req = chunked_req + sched._pending_chunked_abort_req = pending_req + sched.waiting_queue = [] + sched.dllm_config = None + sched.grammar_manager = Mock() + sched.disaggregation_mode = None + sched.enable_hicache_storage = False + sched.mm_receiver = None + sched.ps = SimpleNamespace(pp_size=1) + sched.running_batch = SimpleNamespace(reqs=running_reqs) + sched.last_batch = None + return sched + + +class TestPendingChunkedAbortRace(CustomTestCase): + def test_req_left_chunked_slot_is_aborted(self): + req = _FakeReq("zombie_rid") + sched = _make_scheduler(req, chunked_req=None, running_reqs=[req]) + + sched.process_pending_chunked_abort() + + self.assertIsNotNone(req.to_finish, "recorded abort was never applied") + self.assertIsNone(sched._pending_chunked_abort_req) + + def test_finished_req_only_clears_marker(self): + req = _FakeReq("done_rid") + req._finished = True + sched = _make_scheduler(req, chunked_req=None, running_reqs=[]) + + sched.process_pending_chunked_abort() + + self.assertIsNone(req.to_finish) + self.assertIsNone(sched._pending_chunked_abort_req) + + +if __name__ == "__main__": + unittest.main(verbosity=2) diff --git a/test/registered/unit/managers/test_tokenizer_manager_rid_cleanup.py b/test/registered/unit/managers/test_tokenizer_manager_rid_cleanup.py index 93a3f8bf2..4f54c6e6a 100644 --- a/test/registered/unit/managers/test_tokenizer_manager_rid_cleanup.py +++ b/test/registered/unit/managers/test_tokenizer_manager_rid_cleanup.py @@ -10,11 +10,12 @@ Covers: - _handle_batch_output cleans up rid_to_state on finished requests - _init_req_state rejects duplicate rids - Resubmission succeeds after cleanup + - Handler failures clean up pending and dispatched requests """ import asyncio import unittest -from unittest.mock import AsyncMock, MagicMock, Mock +from unittest.mock import AsyncMock, MagicMock, Mock, patch import msgspec @@ -454,41 +455,88 @@ def _make_generate_obj(rid, is_single): return obj -class TestDiscardPendingReqStates(CustomTestCase): - """Direct tests for _discard_pending_req_states.""" +class TestReleaseReqStatesOnFailure(CustomTestCase): + """Direct tests for _release_req_states_on_failure.""" - def test_discard_single(self): + def test_undelivered_single_is_dropped(self): tm = _make_tokenizer_manager(self) rid = "d_single" tm.rid_to_state[rid] = _make_req_state(rid) - obj = Mock(spec=GenerateReqInput) - obj.is_single = True - obj.rid = rid - tm._discard_pending_req_states(obj) + tm._release_req_states_on_failure([rid]) self.assertNotIn(rid, tm.rid_to_state) - def test_discard_batch_removes_all(self): + def test_undelivered_batch_removes_all(self): tm = _make_tokenizer_manager(self) rids = ["d0", "d1", "d2"] for r in rids: tm.rid_to_state[r] = _make_req_state(r) - obj = Mock(spec=GenerateReqInput) - obj.is_single = False - obj.rid = list(rids) - tm._discard_pending_req_states(obj) + tm._release_req_states_on_failure(rids) for r in rids: self.assertNotIn(r, tm.rid_to_state) - def test_discard_ignores_already_removed(self): - """Popping a rid that is no longer present must not raise.""" + def test_ignores_already_removed(self): + """A rid that is no longer present must not raise.""" tm = _make_tokenizer_manager(self) tm.rid_to_state["p1"] = _make_req_state("p1") - obj = Mock(spec=GenerateReqInput) - obj.is_single = False - obj.rid = ["p1", "already_gone"] - tm._discard_pending_req_states(obj) # must not raise + tm._release_req_states_on_failure(["p1", "already_gone"]) self.assertNotIn("p1", tm.rid_to_state) + def test_dispatched_single_is_aborted_and_state_kept(self): + tm = _make_tokenizer_manager(self) + tm.server_args.tokenizer_worker_num = 1 + tm._dispatch_to_scheduler = Mock() + tm.enable_metrics = True + tm.metrics_collector = MagicMock() + rid = "d_live" + state = _make_req_state(rid) + state.dispatched = True + tm.rid_to_state[rid] = state + tm._release_req_states_on_failure([rid]) + tm._release_req_states_on_failure([rid]) + + sent = [c.args[0] for c in tm._dispatch_to_scheduler.call_args_list] + self.assertEqual( + [type(m) for m in sent], [AbortReq], "expected exactly one AbortReq" + ) + self.assertEqual(sent[0].rid, rid) + self.assertIn(rid, tm.rid_to_state) + self.assertTrue(state.abort_sent) + tm.metrics_collector.observe_one_aborted_request.assert_called_once() + + def test_dispatched_batch_aborts_delivered_and_drops_rest(self): + tm = _make_tokenizer_manager(self) + tm.server_args.tokenizer_worker_num = 1 + tm._dispatch_to_scheduler = Mock() + delivered, undelivered = "d_delivered", "d_undelivered" + live = _make_req_state(delivered) + live.dispatched = True + tm.rid_to_state[delivered] = live + tm.rid_to_state[undelivered] = _make_req_state(undelivered) + tm._release_req_states_on_failure([delivered, undelivered]) + + sent = [c.args[0] for c in tm._dispatch_to_scheduler.call_args_list] + self.assertEqual([type(m) for m in sent], [AbortReq]) + self.assertEqual(sent[0].rid, delivered) + self.assertIn(delivered, tm.rid_to_state) + self.assertNotIn(undelivered, tm.rid_to_state) + + def test_abort_failure_does_not_stop_cleanup(self): + tm = _make_tokenizer_manager(self) + tm.server_args.tokenizer_worker_num = 1 + tm._dispatch_to_scheduler = Mock(side_effect=RuntimeError("send failed")) + delivered, undelivered = "live", "pending" + live = _make_req_state(delivered) + live.dispatched = True + tm.rid_to_state[delivered] = live + tm.rid_to_state[undelivered] = _make_req_state(undelivered) + + with self.assertLogs(level="ERROR"): + tm._release_req_states_on_failure([delivered, undelivered]) + + self.assertIn(delivered, tm.rid_to_state) + self.assertFalse(live.abort_sent) + self.assertNotIn(undelivered, tm.rid_to_state) + class TestParallelStreamTaskCleanup(CustomTestCase): def test_failing_choice_cancels_and_closes_sibling_waiters(self): @@ -595,6 +643,27 @@ class TestGenerateRequestCleanupOnDispatchFailure(CustomTestCase): for r in rids: self.assertNotIn(r, tm.rid_to_state) + def test_parallel_sampling_failure_cleans_generated_rid(self): + tm = _make_tm_for_generate(self) + obj = GenerateReqInput( + text=["hello"], + rid=["base"], + sampling_params={"n": 2}, + ) + tokenized = MagicMock() + tokenized.mm_inputs = None + tokenized.sampling_params = MagicMock() + tm._tokenize_one_request = AsyncMock(return_value=tokenized) + tm._send_one_request = Mock(side_effect=RuntimeError("dispatch failed")) + + async def drive(): + await tm.generate_request(obj).__anext__() + + with self.assertRaisesRegex(RuntimeError, "dispatch failed"): + asyncio.run(drive()) + + self.assertFalse(tm.rid_to_state) + def test_thinking_budget_rejects_runtime_without_strict_thinking(self): tm = _make_tm_for_generate(self) obj = GenerateReqInput( @@ -641,5 +710,52 @@ class TestWaitOneResponseAfterStateFreed(CustomTestCase): self.assertEqual(out["text"], "hello") +class TestDisconnectAfterDispatchAbortsRequest(CustomTestCase): + """Cancellation after dispatch must stop the scheduler request.""" + + @patch( + "sglang.srt.managers.tokenizer_manager.wrap_shm_features", + side_effect=lambda obj: obj, + ) + def test_cancel_after_dispatch_sends_abort_and_keeps_state(self, _wrap_shm): + tm = _make_tm_for_generate(self) + tm.cuda_vmm_feature_transport = Mock() + tm.cuda_vmm_feature_transport.prepare_for_dispatch_async = AsyncMock( + return_value=[] + ) + tm._dispatch_to_scheduler = Mock() + rid = "disconnect_zombie" + obj = _make_generate_obj(rid, is_single=True) + obj.return_prompt_token_ids = False + tokenized = MagicMock() + tokenized.rid = rid + tokenized.mm_inputs = None + tm._tokenize_one_request = AsyncMock(return_value=tokenized) + + async def drive(): + task = asyncio.create_task(tm.generate_request(obj).__anext__()) + for _ in range(100): + await asyncio.sleep(0) + if tm._dispatch_to_scheduler.called: + break + self.assertTrue( + tm._dispatch_to_scheduler.called, "request never dispatched" + ) + state = tm.rid_to_state.get(rid) + self.assertIsNotNone(state) + self.assertTrue(state.dispatched) + + task.cancel() + with self.assertRaises(asyncio.CancelledError): + await task + + asyncio.run(drive()) + + sent = [c.args[0] for c in tm._dispatch_to_scheduler.call_args_list] + aborts = [m for m in sent if isinstance(m, AbortReq) and m.rid == rid] + self.assertTrue(aborts, "disconnect must send an AbortReq to the scheduler") + self.assertIn(rid, tm.rid_to_state) + + if __name__ == "__main__": unittest.main(verbosity=2) diff --git a/test/registered/unit/multimodal/test_gpu_feature_transport.py b/test/registered/unit/multimodal/test_gpu_feature_transport.py index 1c488832f..bef625a72 100644 --- a/test/registered/unit/multimodal/test_gpu_feature_transport.py +++ b/test/registered/unit/multimodal/test_gpu_feature_transport.py @@ -376,6 +376,7 @@ class TestCudaVmmFeatureTransport(unittest.TestCase): from sglang.srt.managers.tokenizer_manager import TokenizerManager manager = object.__new__(TokenizerManager) + manager.rid_to_state = {} transport = MagicMock() transport.prepare_for_dispatch_async = AsyncMock(return_value=[]) manager.cuda_vmm_feature_transport = transport @@ -403,6 +404,7 @@ class TestCudaVmmFeatureTransport(unittest.TestCase): ) manager = object.__new__(tokenizer_manager.TokenizerManager) + manager.rid_to_state = {} transport = MagicMock() manager._dispatch_to_scheduler = MagicMock( side_effect=RuntimeError("send failed") @@ -437,6 +439,7 @@ class TestCudaVmmFeatureTransport(unittest.TestCase): ) manager = object.__new__(tokenizer_manager.TokenizerManager) + manager.rid_to_state = {} transport = MagicMock() manager._dispatch_to_scheduler = MagicMock() time_stats = MagicMock() diff --git a/test/registered/unit/test_chain_read_ratchet.py b/test/registered/unit/test_chain_read_ratchet.py index 99abbcdee..ffd179f3e 100644 --- a/test/registered/unit/test_chain_read_ratchet.py +++ b/test/registered/unit/test_chain_read_ratchet.py @@ -78,9 +78,25 @@ def _returned_field_names(function): returned literal, assignments (annotated or not) to a returned name, a literal-key subscript write on it, and `.update(field=...)` on it. A spelling this cannot see raises instead of skipping. + + ``overrides[name]`` is also accepted when ``name`` comes from + ``for name in ("a", "b", ...)`` -- the keys stay statically enumerable. """ names = set() returned = set() + # for x in ("a", "b"): ... -> {"x": {"a", "b"}} + loop_keys = { + node.target.id: {elt.value for elt in node.iter.elts} + for node in ast.walk(function) + if isinstance(node, ast.For) + and isinstance(node.target, ast.Name) + and isinstance(node.iter, (ast.Tuple, ast.List)) + and node.iter.elts + and all( + isinstance(elt, ast.Constant) and isinstance(elt.value, str) + for elt in node.iter.elts + ) + } def top_level_keys(mapping): for key in mapping.keys: @@ -88,6 +104,14 @@ def _returned_field_names(function): raise AssertionError(f"non-literal key in {function.name}") names.add(key.value) + def add_subscript_key(key): + if isinstance(key, ast.Constant): + names.add(key.value) + elif isinstance(key, ast.Name) and key.id in loop_keys: + names.update(loop_keys[key.id]) + else: + raise AssertionError(f"non-literal key in {function.name}") + for node in ast.walk(function): if isinstance(node, ast.Return) and node.value is not None: value = node.value @@ -114,9 +138,7 @@ def _returned_field_names(function): if isinstance(target, ast.Subscript) and ( isinstance(target.value, ast.Name) and target.value.id in returned ): - if not isinstance(target.slice, ast.Constant): - raise AssertionError(f"non-literal key in {function.name}") - names.add(target.slice.value) + add_subscript_key(target.slice) if ( isinstance(node, ast.Call) and isinstance(node.func, ast.Attribute)