diff --git a/python/sglang/srt/entrypoints/openai/serving_chat.py b/python/sglang/srt/entrypoints/openai/serving_chat.py index 6e50f3278..2f4c6c394 100644 --- a/python/sglang/srt/entrypoints/openai/serving_chat.py +++ b/python/sglang/srt/entrypoints/openai/serving_chat.py @@ -5,6 +5,7 @@ import json import logging import time import uuid +from http import HTTPStatus from typing import TYPE_CHECKING, Any, AsyncGenerator, Dict, List, Optional, Union import jinja2 @@ -641,7 +642,7 @@ class OpenAIServingChat(OpenAIServingBase): routed_experts[index] = content["meta_info"].get("routed_experts", None) # Handle logprobs - finish_reason = content["meta_info"]["finish_reason"] + finish_reason = content["meta_info"].get("finish_reason", None) choice_logprobs = None if request.logprobs: n_prev_token = n_prev_tokens.get(index, 0) @@ -661,7 +662,20 @@ class OpenAIServingChat(OpenAIServingBase): # Track finish_reason for each index if finish_reason_type: - finish_reasons[index] = finish_reason + # If the abort is from scheduler. + if finish_reason_type == "abort": + code = finish_reason.get( + "status_code", HTTPStatus.INTERNAL_SERVER_ERROR + ) + error = self.create_streaming_error_response( + finish_reason.get("message", "Generation aborted."), + code.name, + code.value, + ) + yield f"data: {error}\n\n" + break + else: + finish_reasons[index] = finish_reason # First chunk with role if is_firsts.get(index, True): diff --git a/python/sglang/srt/entrypoints/openai/serving_completions.py b/python/sglang/srt/entrypoints/openai/serving_completions.py index a639d2bed..63b11eca4 100644 --- a/python/sglang/srt/entrypoints/openai/serving_completions.py +++ b/python/sglang/srt/entrypoints/openai/serving_completions.py @@ -2,6 +2,7 @@ from __future__ import annotations import logging import time +from http import HTTPStatus from typing import TYPE_CHECKING, Any, AsyncGenerator, Dict, List, Optional, Union from fastapi import Request @@ -270,13 +271,27 @@ class OpenAIServingCompletion(OpenAIServingBase): # Generate delta delta = text[len(stream_buffer) :] stream_buffers[index] = stream_buffer + delta - finish_reason = content["meta_info"]["finish_reason"] + finish_reason = content["meta_info"].get("finish_reason", None) + finish_reason_type = finish_reason["type"] if finish_reason else None + + # If the abort is from scheduler. + if finish_reason_type == "abort": + code = finish_reason.get( + "status_code", HTTPStatus.INTERNAL_SERVER_ERROR + ) + error = self.create_streaming_error_response( + finish_reason.get("message", "Generation aborted."), + code.name, + code.value, + ) + yield f"data: {error}\n\n" + break choice_data = CompletionResponseStreamChoice( index=index, text=delta, logprobs=logprobs, - finish_reason=finish_reason["type"] if finish_reason else None, + finish_reason=finish_reason_type, matched_stop=( finish_reason["matched"] if finish_reason and "matched" in finish_reason diff --git a/test/registered/openai_server/basic/test_serving_chat.py b/test/registered/openai_server/basic/test_serving_chat.py index feaad53e7..2ad3ec2e7 100644 --- a/test/registered/openai_server/basic/test_serving_chat.py +++ b/test/registered/openai_server/basic/test_serving_chat.py @@ -9,6 +9,7 @@ or import json import unittest import uuid +from http import HTTPStatus from typing import Optional from unittest.mock import Mock, patch @@ -705,6 +706,83 @@ class ServingChatTestCase(unittest.TestCase): serving_chat = OpenAIServingChat(tokenizer_manager, TemplateManager()) self.assertFalse(serving_chat.use_dpsk_v32_encoding) + def test_streaming_abort_yields_error(self): + """Test that an abort finish reason during streaming correctly yields an error and stops.""" + err_msg = "Aborted by scheduler" + err_code = HTTPStatus.INTERNAL_SERVER_ERROR + + async def _mock_generate_abort(): + yield { + "text": "Partial ", + "meta_info": { + "id": "chatcmpl-test", + "prompt_tokens": 10, + "completion_tokens": 2, + "cached_tokens": 0, + "finish_reason": { + "type": "abort", + "status_code": err_code, + "message": err_msg, + }, + "output_token_logprobs": None, + "output_top_logprobs": None, + }, + "index": 0, + } + + self.tm.generate_request.return_value = _mock_generate_abort() + + req = ChatCompletionRequest( + model="x", + messages=[{"role": "user", "content": "Hi?"}], + temperature=0.7, + max_tokens=100, + stream=True, + ) + + with patch( + "sglang.srt.entrypoints.openai.serving_chat.generate_chat_conv" + ) as conv_mock: + # Create a mock conversation object + conv_ins = Mock() + conv_ins.get_prompt.return_value = "Test prompt" + conv_mock.return_value = conv_ins + + adapted_request, _ = self.chat._convert_to_internal_request( + req, self.fastapi_request + ) + + async def run_stream(): + chunks = [] + try: + async for chunk in self.chat._generate_chat_stream( + adapted_request, req, self.fastapi_request + ): + chunks.append(chunk) + except Exception as e: + print(f"Error during stream iteration: {e}") + return chunks + + loop = get_or_create_event_loop() + chunks = loop.run_until_complete(run_stream()) + + error_chunk_data = None + for c in chunks: + if "error" in c: + error_chunk_data = json.loads(c[len("data: ") :]) + break + self.assertIsNotNone(error_chunk_data, "Error chunk not found in stream") + self.assertEqual(error_chunk_data["error"]["message"], err_msg) + self.assertEqual(error_chunk_data["error"]["code"], err_code.value) + + # Ensure the stream stops after the abort error + # The last chunk should be "data: [DONE]\n\n" + self.assertEqual(chunks[-1], "data: [DONE]\n\n") + + # Check that there is an error chunk and a DONE chunk + self.assertEqual(len(chunks), 2) + self.assertIn("error", chunks[0]) + if __name__ == "__main__": unittest.main(verbosity=2) diff --git a/test/registered/openai_server/basic/test_serving_completions.py b/test/registered/openai_server/basic/test_serving_completions.py index 34c262e60..ecbad79d8 100644 --- a/test/registered/openai_server/basic/test_serving_completions.py +++ b/test/registered/openai_server/basic/test_serving_completions.py @@ -4,13 +4,18 @@ Run with: python -m unittest tests.test_serving_completions_unit -v """ +import json import unittest +from http import HTTPStatus from typing import Optional from unittest.mock import AsyncMock, Mock +from fastapi import Request + from sglang.srt.entrypoints.openai.protocol import CompletionRequest from sglang.srt.entrypoints.openai.serving_completions import OpenAIServingCompletion from sglang.srt.managers.tokenizer_manager import TokenizerManager +from sglang.srt.utils import get_or_create_event_loop from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci register_cuda_ci(est_time=10, suite="stage-b-test-small-1-gpu") @@ -49,6 +54,7 @@ class ServingCompletionTestCase(unittest.TestCase): self.template_manager = _MockTemplateManager() self.sc = OpenAIServingCompletion(tm, self.template_manager) + self.fastapi_request = Mock(spec=Request) # ---------- prompt-handling ---------- def test_single_string_prompt(self): @@ -181,6 +187,72 @@ class ServingCompletionTestCase(unittest.TestCase): self.assertEqual(response.choices[0].text, " world") self.assertEqual(len(response.choices[0].logprobs.top_logprobs), 0) + def test_streaming_abort_yields_error(self): + """Test that an abort finish reason during streaming correctly yields an error and stops.""" + err_msg = "Aborted by scheduler" + err_code = HTTPStatus.INTERNAL_SERVER_ERROR + + async def _mock_generate_abort(*args, **kwargs): + yield { + "text": "Partial ", + "meta_info": { + "id": "cmpl-test", + "prompt_tokens": 10, + "completion_tokens": 2, + "cached_tokens": 0, + "finish_reason": { + "type": "abort", + "status_code": err_code, + "message": err_msg, + }, + "output_token_logprobs": None, + "output_top_logprobs": None, + }, + "index": 0, + } + + self.sc.tokenizer_manager.generate_request = _mock_generate_abort + + req = CompletionRequest( + model="x", + prompt="Hello world", + max_tokens=100, + stream=True, + ) + + adapted_request, _ = self.sc._convert_to_internal_request(req) + + async def run_stream(): + chunks = [] + try: + async for chunk in self.sc._generate_completion_stream( + adapted_request, req, self.fastapi_request + ): + chunks.append(chunk) + except Exception as e: + print(f"Error during stream iteration: {e}") + return chunks + + loop = get_or_create_event_loop() + chunks = loop.run_until_complete(run_stream()) + + error_chunk_data = None + for c in chunks: + if "error" in c: + error_chunk_data = json.loads(c[len("data: ") :]) + break + self.assertIsNotNone(error_chunk_data, "Error chunk not found in stream") + self.assertEqual(error_chunk_data["error"]["message"], err_msg) + self.assertEqual(error_chunk_data["error"]["code"], err_code.value) + + # Ensure the stream stops after the abort error + # The last chunk should be "data: [DONE]\n\n" + self.assertEqual(chunks[-1], "data: [DONE]\n\n") + + # Check that there is an error chunk and a DONE chunk, and possibly a role chunk + self.assertGreaterEqual(len(chunks), 2) + self.assertIn("error", chunks[0]) + if __name__ == "__main__": unittest.main(verbosity=2)