Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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."""
Expand Down Expand Up @@ -121,17 +119,16 @@ 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",
400,
("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"),
Comment thread
Shaun-Walsh marked this conversation as resolved.
pytest.param(
"logprobs_top_without_enabled",
400,
Expand Down
40 changes: 40 additions & 0 deletions tests/entrypoints/openai_api/test_api_server_guards.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down
20 changes: 20 additions & 0 deletions vllm_omni/entrypoints/openai/api_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)],
Expand All @@ -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:
Expand Down Expand Up @@ -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)
Expand Down
Loading