diff --git a/python/sglang/srt/entrypoints/openai/serving_chat.py b/python/sglang/srt/entrypoints/openai/serving_chat.py index dcdd60bad..ff09f3e02 100644 --- a/python/sglang/srt/entrypoints/openai/serving_chat.py +++ b/python/sglang/srt/entrypoints/openai/serving_chat.py @@ -77,6 +77,8 @@ if TYPE_CHECKING: logger = logging.getLogger(__name__) +_MEDIA_CONTENT_PART_TYPES = frozenset({"image_url", "video_url", "audio_url"}) + def normalize_tool_content(role: str, content): """Normalize tool message content from OpenAI array format to plain string. @@ -605,6 +607,10 @@ class OpenAIServingChat(OpenAIServingBase): if not request.messages: return "Messages cannot be empty." + media_error = self._validate_media_content(request) + if media_error: + return media_error + if ( isinstance(request.tool_choice, str) and request.tool_choice.lower() == "required" @@ -658,6 +664,28 @@ class OpenAIServingChat(OpenAIServingBase): return None + def _validate_media_content(self, request: ChatCompletionRequest) -> Optional[str]: + if self.tokenizer_manager.model_config.is_multimodal: + return None + + media_type = next( + ( + part.type + for message in request.messages + if isinstance(message.content, list) + for part in message.content + if part.type in _MEDIA_CONTENT_PART_TYPES + ), + None, + ) + if media_type is None: + return None + + return ( + "Model only supports text input; " + f"received unsupported content type '{media_type}'." + ) + def _convert_to_internal_request( self, request: ChatCompletionRequest, diff --git a/python/sglang/srt/entrypoints/openai/serving_responses.py b/python/sglang/srt/entrypoints/openai/serving_responses.py index 66758afaa..58746a909 100644 --- a/python/sglang/srt/entrypoints/openai/serving_responses.py +++ b/python/sglang/srt/entrypoints/openai/serving_responses.py @@ -80,6 +80,10 @@ if TYPE_CHECKING: logger = logging.getLogger(__name__) +class _MediaInputValidationError(ValueError): + pass + + class OpenAIServingResponses(OpenAIServingChat): """Handler for /v1/responses requests""" @@ -239,6 +243,8 @@ class OpenAIServingResponses(OpenAIServingChat): processed_messages, ) = await self._make_request(request, prev_response, tokenizer) + except _MediaInputValidationError as e: + return self.create_error_response(str(e)) except (ValueError, TypeError, RuntimeError, jinja2.TemplateError) as e: logger.exception("Error in preprocessing prompt inputs") return self.create_error_response(f"{e} {e.__cause__}") @@ -481,6 +487,10 @@ class OpenAIServingResponses(OpenAIServingChat): reasoning_effort=(request.reasoning.effort if request.reasoning else None), ) + media_error = self._validate_media_content(chat_request) + if media_error: + raise _MediaInputValidationError(media_error) + is_multimodal = self.tokenizer_manager.model_config.is_multimodal processed_messages = self._process_messages(chat_request, is_multimodal) diff --git a/test/registered/unit/entrypoints/openai/test_serving_chat.py b/test/registered/unit/entrypoints/openai/test_serving_chat.py index 6e2a525e6..236cf3948 100644 --- a/test/registered/unit/entrypoints/openai/test_serving_chat.py +++ b/test/registered/unit/entrypoints/openai/test_serving_chat.py @@ -119,6 +119,74 @@ class ServingChatTestCase(unittest.TestCase): self.fastapi_request = Mock(spec=Request) self.fastapi_request.headers = {} + def test_text_only_model_rejects_media_before_generation(self): + media_parts = { + "image_url": { + "type": "image_url", + "image_url": {"url": "https://example.com/image.png"}, + }, + "video_url": { + "type": "video_url", + "video_url": {"url": "https://example.com/video.mp4"}, + }, + "audio_url": { + "type": "audio_url", + "audio_url": {"url": "https://example.com/audio.wav"}, + }, + } + + for media_type, media_part in media_parts.items(): + with self.subTest(media_type=media_type): + request = ChatCompletionRequest( + model="x", + messages=[ + { + "role": "user", + "content": [ + {"type": "text", "text": "describe"}, + media_part, + ], + } + ], + ) + response = get_or_create_event_loop().run_until_complete( + self.chat.handle_request(request, self.fastapi_request) + ) + error = json.loads(response.body) + self.assertEqual(response.status_code, HTTPStatus.BAD_REQUEST) + self.assertEqual(error["type"], "BadRequestError") + self.assertIn(media_type, error["message"]) + self.tm.generate_request.assert_not_called() + + def test_media_validation_does_not_reject_supported_content(self): + text_request = ChatCompletionRequest( + model="x", + messages=[ + { + "role": "user", + "content": [{"type": "tool_reference", "name": "get_weather"}], + } + ], + ) + self.assertIsNone(self.chat._validate_request(text_request)) + + self.tm.model_config.is_multimodal = True + multimodal_request = ChatCompletionRequest( + model="x", + messages=[ + { + "role": "user", + "content": [ + { + "type": "image_url", + "image_url": {"url": "https://example.com/image.png"}, + } + ], + } + ], + ) + self.assertIsNone(self.chat._validate_request(multimodal_request)) + # ------------- conversion tests ------------- def test_convert_to_internal_request_single(self): with ( diff --git a/test/registered/unit/entrypoints/openai/test_serving_responses.py b/test/registered/unit/entrypoints/openai/test_serving_responses.py index 61afb21e6..bba25a877 100644 --- a/test/registered/unit/entrypoints/openai/test_serving_responses.py +++ b/test/registered/unit/entrypoints/openai/test_serving_responses.py @@ -303,6 +303,33 @@ class FullResponseUsageTestCase(unittest.TestCase): class MultimodalRequestTestCase(unittest.TestCase): + def test_text_only_create_responses_rejects_media_before_generation(self): + serving = make_serving() + serving._process_messages = Mock() + request = ResponsesRequest( + model="x", + input=[ + { + "role": "user", + "content": [ + {"type": "input_text", "text": "describe it"}, + { + "type": "input_image", + "image_url": "http://example.com/cat.png", + }, + ], + } + ], + store=False, + ) + + response = asyncio.run(serving.create_responses(request)) + + self.assertEqual(response.status_code, 400) + self.assertIn(b"received unsupported content type 'image_url'", response.body) + serving._process_messages.assert_not_called() + serving.tokenizer_manager.generate_request.assert_not_called() + def test_multimodal_create_responses_sends_text_and_media_to_engine(self): serving = make_serving(is_multimodal=True) captured = {}