fix: potential crash for missing stream attribute (#15644)
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user