Return HTTP 400 for streaming validation errors (#21900)
This commit is contained in:
@@ -602,10 +602,25 @@ class OpenAIServingChat(OpenAIServingBase):
|
||||
adapted_request: GenerateReqInput,
|
||||
request: ChatCompletionRequest,
|
||||
raw_request: Request,
|
||||
) -> StreamingResponse:
|
||||
) -> Union[StreamingResponse, ErrorResponse]:
|
||||
"""Handle streaming chat completion request"""
|
||||
generator = self._generate_chat_stream(adapted_request, request, raw_request)
|
||||
|
||||
# Kick-start the generator to trigger validation before HTTP 200 is sent.
|
||||
# If validation fails (e.g., context length exceeded), we can still return
|
||||
# a proper HTTP 400 error response instead of streaming it as SSE payload.
|
||||
try:
|
||||
first_chunk = await generator.__anext__()
|
||||
except ValueError as e:
|
||||
return self.create_error_response(str(e))
|
||||
|
||||
async def prepend_first_chunk():
|
||||
yield first_chunk
|
||||
async for chunk in generator:
|
||||
yield chunk
|
||||
|
||||
return StreamingResponse(
|
||||
self._generate_chat_stream(adapted_request, request, raw_request),
|
||||
prepend_first_chunk(),
|
||||
media_type="text/event-stream",
|
||||
background=self.tokenizer_manager.create_abort_task(adapted_request),
|
||||
)
|
||||
@@ -635,6 +650,7 @@ class OpenAIServingChat(OpenAIServingBase):
|
||||
hidden_states = {}
|
||||
routed_experts = {}
|
||||
|
||||
stream_started = False
|
||||
try:
|
||||
async for content in self.tokenizer_manager.generate_request(
|
||||
adapted_request, raw_request
|
||||
@@ -699,6 +715,7 @@ class OpenAIServingChat(OpenAIServingBase):
|
||||
model=request.model,
|
||||
)
|
||||
yield f"data: {chunk.model_dump_json()}\n\n"
|
||||
stream_started = True
|
||||
|
||||
stream_buffer = stream_buffers.get(index, "")
|
||||
delta = content["text"][len(stream_buffer) :]
|
||||
@@ -879,6 +896,8 @@ class OpenAIServingChat(OpenAIServingBase):
|
||||
yield f"data: {usage_chunk.model_dump_json()}\n\n"
|
||||
|
||||
except ValueError as e:
|
||||
if not stream_started:
|
||||
raise
|
||||
error = self.create_streaming_error_response(str(e))
|
||||
yield f"data: {error}\n\n"
|
||||
|
||||
|
||||
@@ -180,10 +180,25 @@ class OpenAIServingCompletion(OpenAIServingBase):
|
||||
adapted_request: GenerateReqInput,
|
||||
request: CompletionRequest,
|
||||
raw_request: Request,
|
||||
) -> StreamingResponse:
|
||||
) -> Union[StreamingResponse, ErrorResponse]:
|
||||
"""Handle streaming completion request"""
|
||||
generator = self._generate_completion_stream(
|
||||
adapted_request, request, raw_request
|
||||
)
|
||||
|
||||
# Kick-start the generator to trigger validation before HTTP 200 is sent.
|
||||
try:
|
||||
first_chunk = await generator.__anext__()
|
||||
except ValueError as e:
|
||||
return self.create_error_response(str(e))
|
||||
|
||||
async def prepend_first_chunk():
|
||||
yield first_chunk
|
||||
async for chunk in generator:
|
||||
yield chunk
|
||||
|
||||
return StreamingResponse(
|
||||
self._generate_completion_stream(adapted_request, request, raw_request),
|
||||
prepend_first_chunk(),
|
||||
media_type="text/event-stream",
|
||||
background=self.tokenizer_manager.create_abort_task(adapted_request),
|
||||
)
|
||||
@@ -208,6 +223,7 @@ class OpenAIServingCompletion(OpenAIServingBase):
|
||||
hidden_states = {}
|
||||
routed_experts = {}
|
||||
|
||||
stream_started = False
|
||||
try:
|
||||
async for content in self.tokenizer_manager.generate_request(
|
||||
adapted_request, raw_request
|
||||
@@ -312,6 +328,7 @@ class OpenAIServingCompletion(OpenAIServingBase):
|
||||
)
|
||||
|
||||
yield f"data: {chunk.model_dump_json()}\n\n"
|
||||
stream_started = True
|
||||
|
||||
if request.return_hidden_states and hidden_states:
|
||||
for index, choice_hidden_states in hidden_states.items():
|
||||
@@ -373,6 +390,8 @@ class OpenAIServingCompletion(OpenAIServingBase):
|
||||
yield f"data: {final_usage_data}\n\n"
|
||||
|
||||
except Exception as e:
|
||||
if not stream_started:
|
||||
raise
|
||||
error = self.create_streaming_error_response(str(e))
|
||||
yield f"data: {error}\n\n"
|
||||
|
||||
|
||||
@@ -67,6 +67,23 @@ class TestRequestLengthValidation(CustomTestCase):
|
||||
|
||||
self.assertIn("is longer than the model's context length", str(cm.exception))
|
||||
|
||||
def test_input_length_longer_than_context_length_streaming(self):
|
||||
client = openai.Client(api_key=self.api_key, base_url=f"{self.base_url}/v1")
|
||||
|
||||
long_text = "hello " * 1200
|
||||
|
||||
with self.assertRaises(openai.BadRequestError) as cm:
|
||||
client.chat.completions.create(
|
||||
model=DEFAULT_SMALL_MODEL_NAME_FOR_TEST,
|
||||
messages=[
|
||||
{"role": "user", "content": long_text},
|
||||
],
|
||||
temperature=0,
|
||||
stream=True,
|
||||
)
|
||||
|
||||
self.assertIn("is longer than the model's context length", str(cm.exception))
|
||||
|
||||
def test_max_tokens_validation(self):
|
||||
client = openai.Client(api_key=self.api_key, base_url=f"{self.base_url}/v1")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user