fix: potential crash for missing stream attribute (#15644)

This commit is contained in:
shuwenn
2025-12-23 23:24:31 +08:00
committed by GitHub
parent 6a5764a719
commit 53f974b973

View File

@@ -1014,6 +1014,8 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
request: Optional[fastapi.Request] = None,
):
"""Wait for the response of one request."""
# Not all request types have `stream` (e.g., EmbeddingReqInput). Default to non-streaming.
is_stream = getattr(obj, "stream", False)
while True:
try:
await asyncio.wait_for(state.event.wait(), timeout=4)
@@ -1063,7 +1065,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
finish_reason.get("type") == "abort"
and finish_reason.get("status_code") == HTTPStatus.BAD_REQUEST
):
if not obj.stream:
if not is_stream:
raise ValueError(finish_reason["message"])
else:
yield out
@@ -1084,7 +1086,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
# Mark ongoing LoRA request as finished.
if self.server_args.enable_lora and state.obj.lora_path:
await self.lora_registry.release(state.obj.lora_id)
if not obj.stream:
if not is_stream:
raise fastapi.HTTPException(
status_code=finish_reason["status_code"],
detail=finish_reason["message"],
@@ -1097,7 +1099,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
state.event.clear()
if obj.stream:
if is_stream:
# Record response sent time right before we send response.
if not state.response_sent_to_client_ts:
state.response_sent_to_client_ts = time.time()
@@ -1600,7 +1602,9 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
if isinstance(recv_obj, BatchStrOutput):
state.text += recv_obj.output_strs[i]
if self.server_args.stream_output and state.obj.stream:
# Not all request types have `stream` (e.g., EmbeddingReqInput). Default to non-streaming.
is_stream = getattr(state.obj, "stream", False)
if self.server_args.stream_output and is_stream:
state.output_ids.extend(recv_obj.output_ids[i])
output_token_ids = state.output_ids[state.last_output_offset :]
state.last_output_offset = len(state.output_ids)
@@ -1614,7 +1618,8 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
"meta_info": meta_info,
}
elif isinstance(recv_obj, BatchTokenIDOutput):
if self.server_args.stream_output and state.obj.stream:
is_stream = getattr(state.obj, "stream", False)
if self.server_args.stream_output and is_stream:
state.output_ids.extend(recv_obj.output_ids[i])
output_token_ids = state.output_ids[state.last_output_offset :]
state.last_output_offset = len(state.output_ids)