diff --git a/tests/dfx/reliability/invalid_param_test/test_invalid_omni_chat.py b/tests/dfx/reliability/invalid_param_test/test_invalid_omni_chat.py index 835d546e485..a6d5fd342db 100644 --- a/tests/dfx/reliability/invalid_param_test/test_invalid_omni_chat.py +++ b/tests/dfx/reliability/invalid_param_test/test_invalid_omni_chat.py @@ -15,8 +15,6 @@ pytestmark = [pytest.mark.slow, pytest.mark.omni] -_SKIP_ISSUE_3649 = pytest.mark.skip(reason="https://github.com/vllm-project/vllm-omni/issues/3649") - def _minimal_chat_json(omni_server: OmniServer) -> dict[str, object]: """Minimal valid chat body; individual tests override one offending field.""" @@ -121,9 +119,8 @@ def _chat_completions_request_without_expectations(omni_server: OmniServer, case pytest.param( "modalities_list_bad", 400, - ("modalities", "value_error", ""), + ("modalities", "list of strings"), id="modalities_list_bad_element", - marks=_SKIP_ISSUE_3649, ), pytest.param( "response_format_json_schema_incomplete", @@ -131,7 +128,7 @@ def _chat_completions_request_without_expectations(omni_server: OmniServer, case ("response_format", "BadRequestError", "json_schema"), id="invalid_response_format_json_schema", ), - pytest.param("logprobs_wrong_type", 400, "logprobs", id="logprobs_wrong_type", marks=_SKIP_ISSUE_3649), + pytest.param("logprobs_wrong_type", 400, ("logprobs", "boolean"), id="logprobs_wrong_type"), pytest.param( "logprobs_top_without_enabled", 400, diff --git a/tests/entrypoints/openai_api/test_api_server_guards.py b/tests/entrypoints/openai_api/test_api_server_guards.py index 43d88d6925c..58700b55207 100644 --- a/tests/entrypoints/openai_api/test_api_server_guards.py +++ b/tests/entrypoints/openai_api/test_api_server_guards.py @@ -821,6 +821,46 @@ def test_speech_without_handler_preserves_not_found_http_error() -> None: assert exc_info.value.detail == "The model does not support Speech API" +@pytest.mark.parametrize( + ("field", "value", "detail"), + [ + ("modalities", [123], "modalities must be a list of strings"), + ("logprobs", "yes", "logprobs must be a boolean"), + ], +) +def test_chat_completion_raw_body_guards_reject_lax_types(field, value, detail) -> None: + """The HTTP boundary rejects values that upstream Pydantic would coerce.""" + app = FastAPI() + app.state.openai_serving_chat = None + app.state.serving_tokenization = None + app.add_api_route("/v1/chat/completions", api_server.create_chat_completion, methods=["POST"]) + client = TestClient(app) + + payload = { + "model": "demo-model", + "messages": [{"role": "user", "content": "hello"}], + "stream": False, + field: value, + } + response = client.post("/v1/chat/completions", json=payload) + + assert response.status_code == 400 + assert detail in response.json()["detail"] + + +@pytest.mark.parametrize( + "raw_body", + [ + {"modalities": None}, + {"logprobs": None}, + {"modalities": None, "logprobs": None}, + ], +) +def test_chat_completion_raw_body_guards_allow_null_defaults(raw_body) -> None: + """Explicit JSON null keeps the upstream request model's default behavior.""" + api_server._validate_chat_completion_raw_body(raw_body) + + @pytest.mark.asyncio async def test_multi_api_rejects_runtime_voice_upload() -> None: app = FastAPI() diff --git a/vllm_omni/entrypoints/openai/api_server.py b/vllm_omni/entrypoints/openai/api_server.py index 5f3a4bf182d..a1698e379b8 100644 --- a/vllm_omni/entrypoints/openai/api_server.py +++ b/vllm_omni/entrypoints/openai/api_server.py @@ -1198,6 +1198,22 @@ async def omni_init_app_state( state.server_load_metrics = 0 +def _validate_chat_completion_raw_body(raw_body: dict[str, Any]) -> None: + """Reject values that upstream ChatCompletionRequest coerces too broadly.""" + if "modalities" in raw_body and raw_body["modalities"] is not None: + modalities = raw_body["modalities"] + if not isinstance(modalities, list) or not all(isinstance(m, str) for m in modalities): + raise HTTPException( + status_code=HTTPStatus.BAD_REQUEST.value, + detail='modalities must be a list of strings, e.g. ["text", "audio", "image"]', + ) + if "logprobs" in raw_body and raw_body["logprobs"] is not None and not isinstance(raw_body["logprobs"], bool): + raise HTTPException( + status_code=HTTPStatus.BAD_REQUEST.value, + detail="logprobs must be a boolean (true or false)", + ) + + @router.post( "/v1/chat/completions", dependencies=[Depends(validate_json_request)], @@ -1211,6 +1227,8 @@ async def omni_init_app_state( @with_cancellation @load_aware_call async def create_chat_completion(request: ChatCompletionRequest, raw_request: Request): + raw_body = await raw_request.json() + _validate_chat_completion_raw_body(raw_body) metrics_header_format = raw_request.headers.get(ENDPOINT_LOAD_METRICS_FORMAT_HEADER_LABEL, "") handler = Omnichat(raw_request) if handler is None: @@ -1287,6 +1305,8 @@ async def create_chat_completion(request: ChatCompletionRequest, raw_request: Re @with_cancellation @load_aware_call async def create_batch_chat_completion(request: BatchChatCompletionRequest, raw_request: Request): + raw_body = await raw_request.json() + _validate_chat_completion_raw_body(raw_body) handler = OmniBatchChat(raw_request) if handler is None: base_server = getattr(raw_request.app.state, "serving_tokenization", None)