Return HTTP 400 for streaming validation errors (#21900)
This commit is contained in:
@@ -602,10 +602,25 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
adapted_request: GenerateReqInput,
|
adapted_request: GenerateReqInput,
|
||||||
request: ChatCompletionRequest,
|
request: ChatCompletionRequest,
|
||||||
raw_request: Request,
|
raw_request: Request,
|
||||||
) -> StreamingResponse:
|
) -> Union[StreamingResponse, ErrorResponse]:
|
||||||
"""Handle streaming chat completion request"""
|
"""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(
|
return StreamingResponse(
|
||||||
self._generate_chat_stream(adapted_request, request, raw_request),
|
prepend_first_chunk(),
|
||||||
media_type="text/event-stream",
|
media_type="text/event-stream",
|
||||||
background=self.tokenizer_manager.create_abort_task(adapted_request),
|
background=self.tokenizer_manager.create_abort_task(adapted_request),
|
||||||
)
|
)
|
||||||
@@ -635,6 +650,7 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
hidden_states = {}
|
hidden_states = {}
|
||||||
routed_experts = {}
|
routed_experts = {}
|
||||||
|
|
||||||
|
stream_started = False
|
||||||
try:
|
try:
|
||||||
async for content in self.tokenizer_manager.generate_request(
|
async for content in self.tokenizer_manager.generate_request(
|
||||||
adapted_request, raw_request
|
adapted_request, raw_request
|
||||||
@@ -699,6 +715,7 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
model=request.model,
|
model=request.model,
|
||||||
)
|
)
|
||||||
yield f"data: {chunk.model_dump_json()}\n\n"
|
yield f"data: {chunk.model_dump_json()}\n\n"
|
||||||
|
stream_started = True
|
||||||
|
|
||||||
stream_buffer = stream_buffers.get(index, "")
|
stream_buffer = stream_buffers.get(index, "")
|
||||||
delta = content["text"][len(stream_buffer) :]
|
delta = content["text"][len(stream_buffer) :]
|
||||||
@@ -879,6 +896,8 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
yield f"data: {usage_chunk.model_dump_json()}\n\n"
|
yield f"data: {usage_chunk.model_dump_json()}\n\n"
|
||||||
|
|
||||||
except ValueError as e:
|
except ValueError as e:
|
||||||
|
if not stream_started:
|
||||||
|
raise
|
||||||
error = self.create_streaming_error_response(str(e))
|
error = self.create_streaming_error_response(str(e))
|
||||||
yield f"data: {error}\n\n"
|
yield f"data: {error}\n\n"
|
||||||
|
|
||||||
|
|||||||
@@ -180,10 +180,25 @@ class OpenAIServingCompletion(OpenAIServingBase):
|
|||||||
adapted_request: GenerateReqInput,
|
adapted_request: GenerateReqInput,
|
||||||
request: CompletionRequest,
|
request: CompletionRequest,
|
||||||
raw_request: Request,
|
raw_request: Request,
|
||||||
) -> StreamingResponse:
|
) -> Union[StreamingResponse, ErrorResponse]:
|
||||||
"""Handle streaming completion request"""
|
"""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(
|
return StreamingResponse(
|
||||||
self._generate_completion_stream(adapted_request, request, raw_request),
|
prepend_first_chunk(),
|
||||||
media_type="text/event-stream",
|
media_type="text/event-stream",
|
||||||
background=self.tokenizer_manager.create_abort_task(adapted_request),
|
background=self.tokenizer_manager.create_abort_task(adapted_request),
|
||||||
)
|
)
|
||||||
@@ -208,6 +223,7 @@ class OpenAIServingCompletion(OpenAIServingBase):
|
|||||||
hidden_states = {}
|
hidden_states = {}
|
||||||
routed_experts = {}
|
routed_experts = {}
|
||||||
|
|
||||||
|
stream_started = False
|
||||||
try:
|
try:
|
||||||
async for content in self.tokenizer_manager.generate_request(
|
async for content in self.tokenizer_manager.generate_request(
|
||||||
adapted_request, raw_request
|
adapted_request, raw_request
|
||||||
@@ -312,6 +328,7 @@ class OpenAIServingCompletion(OpenAIServingBase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
yield f"data: {chunk.model_dump_json()}\n\n"
|
yield f"data: {chunk.model_dump_json()}\n\n"
|
||||||
|
stream_started = True
|
||||||
|
|
||||||
if request.return_hidden_states and hidden_states:
|
if request.return_hidden_states and hidden_states:
|
||||||
for index, choice_hidden_states in hidden_states.items():
|
for index, choice_hidden_states in hidden_states.items():
|
||||||
@@ -373,6 +390,8 @@ class OpenAIServingCompletion(OpenAIServingBase):
|
|||||||
yield f"data: {final_usage_data}\n\n"
|
yield f"data: {final_usage_data}\n\n"
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
|
if not stream_started:
|
||||||
|
raise
|
||||||
error = self.create_streaming_error_response(str(e))
|
error = self.create_streaming_error_response(str(e))
|
||||||
yield f"data: {error}\n\n"
|
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))
|
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):
|
def test_max_tokens_validation(self):
|
||||||
client = openai.Client(api_key=self.api_key, base_url=f"{self.base_url}/v1")
|
client = openai.Client(api_key=self.api_key, base_url=f"{self.base_url}/v1")
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user