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