diff --git a/python/sglang/benchmark/one_batch_server.py b/python/sglang/benchmark/one_batch_server.py index de95aeb7e..e41fdb249 100644 --- a/python/sglang/benchmark/one_batch_server.py +++ b/python/sglang/benchmark/one_batch_server.py @@ -531,6 +531,10 @@ def run_one_case( gsp_system_prompt_len=gsp_system_prompt_len, gsp_question_len=gsp_question_len, gsp_output_len=gsp_output_len, + # The generated-shared-prefix dataset's from_args requires these; the + # batch-bench path only ever uses the uniform group distribution. + gsp_group_distribution="uniform", + gsp_zipf_alpha=None, ) tok_inner = getattr(tokenizer, "tokenizer", tokenizer) dataset_model_id = model_name or getattr(tok_inner, "name_or_path", None) diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index 5312e9c3c..f9e6cb932 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -597,30 +597,41 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): if self.server_args.tokenizer_worker_num > 1: self._attach_multi_http_worker_info(obj) self._init_req_state(obj, request) - if self.server_args.language_only: - self._handle_epd_disaggregation_encode_request(obj) + try: + if self.server_args.language_only: + self._handle_epd_disaggregation_encode_request(obj) - # Log the request - self.request_logger.log_received_request(obj, self.tokenizer, request) + # Log the request + self.request_logger.log_received_request(obj, self.tokenizer, request) - async with self.is_pause_cond: - await self.is_pause_cond.wait_for(lambda: not self.is_pause) + async with self.is_pause_cond: + await self.is_pause_cond.wait_for(lambda: not self.is_pause) - async with self.model_update_lock.reader_lock: - await self._validate_and_resolve_lora(obj) + async with self.model_update_lock.reader_lock: + await self._validate_and_resolve_lora(obj) - # Tokenize the request and send it to the scheduler - if obj.is_single: - tokenized_obj = await self._tokenize_one_request(obj) - state = self.rid_to_state[obj.rid] - if obj.return_prompt_token_ids: - state.prompt_token_ids = list(tokenized_obj.input_ids) - self._send_one_request(tokenized_obj) - async for response in self._wait_one_response(obj, request): - yield response - else: - async for response in self._handle_batch_request(obj, request): - yield response + # Tokenize the request and send it to the scheduler + if obj.is_single: + tokenized_obj = await self._tokenize_one_request(obj) + state = self.rid_to_state[obj.rid] + if obj.return_prompt_token_ids: + state.prompt_token_ids = list(tokenized_obj.input_ids) + self._send_one_request(tokenized_obj) + async for response in self._wait_one_response(obj, request): + yield response + else: + async for response in self._handle_batch_request(obj, request): + yield response + except Exception: + # _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) + raise def _detect_input_format( self, texts: Union[str, List[str]], is_cross_encoder: bool @@ -2838,6 +2849,20 @@ 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*. + + 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. + """ + 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) + def _should_dispatch_to_encoder( self, obj: Union[GenerateReqInput, EmbeddingReqInput] ) -> bool: 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 95a8ffffd..bfbd6ec41 100644 --- a/test/registered/unit/managers/test_tokenizer_manager_rid_cleanup.py +++ b/test/registered/unit/managers/test_tokenizer_manager_rid_cleanup.py @@ -15,7 +15,7 @@ Covers: import asyncio import dataclasses import unittest -from unittest.mock import MagicMock, Mock +from unittest.mock import AsyncMock, MagicMock, Mock from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.test_utils import CustomTestCase, maybe_stub_sgl_kernel @@ -386,5 +386,134 @@ class TestResubmitAfterCompletion(CustomTestCase): self.assertIn(rid, tm.rid_to_state) +class _DummyAsyncCM: + """Reusable no-op async context manager (stands in for an RW lock).""" + + async def __aenter__(self): + return self + + async def __aexit__(self, *exc): + return False + + +def _make_tm_for_generate() -> TokenizerManager: + """Augment the mocked TokenizerManager with what generate_request needs.""" + tm = _make_tokenizer_manager() + tm.server_args.language_only = False + tm.server_args.tokenizer_worker_num = 1 + tm.auto_create_handle_loop = Mock() + tm._set_default_priority = Mock() + tm.request_logger = Mock() + tm.tokenizer = None + tm.is_pause = False + tm.is_pause_cond = asyncio.Condition() + tm.model_update_lock = Mock() + tm.model_update_lock.reader_lock = _DummyAsyncCM() + tm._validate_and_resolve_lora = AsyncMock(return_value=None) + return tm + + +def _make_generate_obj(rid, is_single): + obj = MagicMock(spec=GenerateReqInput) + obj.routed_dp_rank = None + obj.is_single = is_single + obj.rid = rid + obj.received_time = 0.0 + obj.external_trace_header = None + obj.bootstrap_room = None + obj.normalize_batch_and_arguments = Mock() + if not is_single: + obj.__getitem__.side_effect = lambda i: Mock() + return obj + + +class TestDiscardPendingReqStates(CustomTestCase): + """Direct tests for _discard_pending_req_states.""" + + def test_discard_single(self): + tm = _make_tokenizer_manager() + 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) + self.assertNotIn(rid, tm.rid_to_state) + + def test_discard_batch_removes_all(self): + tm = _make_tokenizer_manager() + 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) + 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.""" + tm = _make_tokenizer_manager() + 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 + self.assertNotIn("p1", tm.rid_to_state) + + +class TestGenerateRequestCleanupOnDispatchFailure(CustomTestCase): + """generate_request must not leak rid_to_state when dispatch fails. + + Regression guard: _init_req_state creates rid_to_state entries up front, + and the only remover is the scheduler-response path. A failure before the + request reaches the scheduler (e.g. input-length validation rejecting an + over-context request) used to leak those entries permanently. + """ + + def test_single_failure_before_dispatch_cleans_up(self): + tm = _make_tm_for_generate() + rid = "single_overlen" + obj = _make_generate_obj(rid, is_single=True) + # Simulate over-length rejection during tokenization/validation. + tm._tokenize_one_request = AsyncMock(side_effect=ValueError("input too long")) + tm._send_one_request = Mock() + + async def drive(): + await tm.generate_request(obj).__anext__() + + with self.assertRaises(ValueError): + asyncio.run(drive()) + + # Got past _init_req_state (which created the entry) ... + tm._tokenize_one_request.assert_awaited_once() + tm._send_one_request.assert_not_called() + # ... and the entry was cleaned up rather than leaked. + self.assertNotIn(rid, tm.rid_to_state) + + def test_batch_failure_before_dispatch_cleans_up_all(self): + tm = _make_tm_for_generate() + rids = ["b0", "b1", "b2"] + obj = _make_generate_obj(list(rids), is_single=False) + + # One over-length sub-request makes the whole batch dispatch raise. + async def _boom(*args, **kwargs): + raise ValueError("input too long") + yield # pragma: no cover (marks this an async generator) + + tm._handle_batch_request = _boom + + async def drive(): + await tm.generate_request(obj).__anext__() + + with self.assertRaises(ValueError): + asyncio.run(drive()) + + # All sub-request entries created by _init_req_state are cleaned up. + for r in rids: + self.assertNotIn(r, tm.rid_to_state) + + if __name__ == "__main__": unittest.main(verbosity=2)