Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 11 additions & 2 deletions python/sglang/srt/managers/tokenizer_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -1688,13 +1688,22 @@ async def _handle_abort_finish_reason(

return None

async def _wait_one_response(
def _wait_one_response(
self,
obj: Union[GenerateReqInput, EmbeddingReqInput],
request: Optional[fastapi.Request] = None,
):
"""Wait for the response of one request."""
# Batch dispatch builds every waiter before advancing any.
# Both removers append the output after the del, so the ReqState stays valid.
state = self.rid_to_state[obj.rid]
return self._stream_one_response(obj=obj, state=state, request=request)

async def _stream_one_response(
self,
obj: Union[GenerateReqInput, EmbeddingReqInput],
state: ReqState,
request: Optional[fastapi.Request] = None,
):
# Not all request types have `stream` (e.g., EmbeddingReqInput). Default to non-streaming.
is_stream = getattr(obj, "stream", False)
while True:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -613,5 +613,33 @@ async def drive():
self.assertFalse(tm.rid_to_state)


class TestWaitOneResponseAfterStateFreed(CustomTestCase):
"""A waiter built before its request finishes must still deliver the output.

Batch dispatch builds every waiter before advancing any, and the
scheduler-response path drops rid_to_state as soon as a request finishes.
"""

def test_generator_built_before_finish_still_delivers_output(self):
tm = _make_tokenizer_manager(self)
tm.request_logger = Mock()
tm.request_metrics_exporter_manager = MagicMock()
tm.request_metrics_exporter_manager.exporter_enabled.return_value = False
rid = "freed_state_rid"
state = _make_req_state(rid)
state.obj.background = True # skip the fastapi disconnect probe
tm.rid_to_state[rid] = state

async def drive():
waiter = tm._wait_one_response(state.obj, None)
await tm._handle_batch_output(_make_batch_str_output(rid))
self.assertNotIn(rid, tm.rid_to_state)
return await waiter.__anext__()

out = asyncio.run(drive())
self.assertEqual(out["meta_info"]["id"], rid)
self.assertEqual(out["text"], "hello")


if __name__ == "__main__":
unittest.main(verbosity=2)
Loading