[fix]reject media input for text-only models (#32914)
Co-authored-by: Xinyuan Tong <xinyuantong.cs@gmail.com>
This commit is contained in:
co-authored by
Xinyuan Tong
parent
e3d4f48e55
commit
690de097c4
@@ -77,6 +77,8 @@ if TYPE_CHECKING:
|
|||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
_MEDIA_CONTENT_PART_TYPES = frozenset({"image_url", "video_url", "audio_url"})
|
||||||
|
|
||||||
|
|
||||||
def normalize_tool_content(role: str, content):
|
def normalize_tool_content(role: str, content):
|
||||||
"""Normalize tool message content from OpenAI array format to plain string.
|
"""Normalize tool message content from OpenAI array format to plain string.
|
||||||
@@ -605,6 +607,10 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
if not request.messages:
|
if not request.messages:
|
||||||
return "Messages cannot be empty."
|
return "Messages cannot be empty."
|
||||||
|
|
||||||
|
media_error = self._validate_media_content(request)
|
||||||
|
if media_error:
|
||||||
|
return media_error
|
||||||
|
|
||||||
if (
|
if (
|
||||||
isinstance(request.tool_choice, str)
|
isinstance(request.tool_choice, str)
|
||||||
and request.tool_choice.lower() == "required"
|
and request.tool_choice.lower() == "required"
|
||||||
@@ -658,6 +664,28 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
|
|
||||||
return None
|
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(
|
def _convert_to_internal_request(
|
||||||
self,
|
self,
|
||||||
request: ChatCompletionRequest,
|
request: ChatCompletionRequest,
|
||||||
|
|||||||
@@ -80,6 +80,10 @@ if TYPE_CHECKING:
|
|||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class _MediaInputValidationError(ValueError):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
class OpenAIServingResponses(OpenAIServingChat):
|
class OpenAIServingResponses(OpenAIServingChat):
|
||||||
"""Handler for /v1/responses requests"""
|
"""Handler for /v1/responses requests"""
|
||||||
|
|
||||||
@@ -239,6 +243,8 @@ class OpenAIServingResponses(OpenAIServingChat):
|
|||||||
processed_messages,
|
processed_messages,
|
||||||
) = await self._make_request(request, prev_response, tokenizer)
|
) = 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:
|
except (ValueError, TypeError, RuntimeError, jinja2.TemplateError) as e:
|
||||||
logger.exception("Error in preprocessing prompt inputs")
|
logger.exception("Error in preprocessing prompt inputs")
|
||||||
return self.create_error_response(f"{e} {e.__cause__}")
|
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),
|
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
|
is_multimodal = self.tokenizer_manager.model_config.is_multimodal
|
||||||
processed_messages = self._process_messages(chat_request, is_multimodal)
|
processed_messages = self._process_messages(chat_request, is_multimodal)
|
||||||
|
|
||||||
|
|||||||
@@ -119,6 +119,74 @@ class ServingChatTestCase(unittest.TestCase):
|
|||||||
self.fastapi_request = Mock(spec=Request)
|
self.fastapi_request = Mock(spec=Request)
|
||||||
self.fastapi_request.headers = {}
|
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 -------------
|
# ------------- conversion tests -------------
|
||||||
def test_convert_to_internal_request_single(self):
|
def test_convert_to_internal_request_single(self):
|
||||||
with (
|
with (
|
||||||
|
|||||||
@@ -303,6 +303,33 @@ class FullResponseUsageTestCase(unittest.TestCase):
|
|||||||
|
|
||||||
|
|
||||||
class MultimodalRequestTestCase(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):
|
def test_multimodal_create_responses_sends_text_and_media_to_engine(self):
|
||||||
serving = make_serving(is_multimodal=True)
|
serving = make_serving(is_multimodal=True)
|
||||||
captured = {}
|
captured = {}
|
||||||
|
|||||||
Reference in New Issue
Block a user