Return HTTP 400 for streaming validation errors (#21900)

This commit is contained in:
Liangsheng Yin
2026-04-01 21:58:12 -07:00
committed by GitHub
parent 153359b4dd
commit 269589ad71
3 changed files with 59 additions and 4 deletions
@@ -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")