fix(security): add resource bounds validation to derender endpoints (#47260)

Signed-off-by: jperezde <jperezde@redhat.com>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
This commit is contained in:
Juan Pérez de Algaba
2026-07-07 14:58:26 +08:00
committed by GitHub
co-authored by mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
parent e040899a00
commit 8e61b646e2
3 changed files with 272 additions and 3 deletions
@@ -489,6 +489,200 @@ async def test_derender_completion_kv_transfer_params_passthrough(client):
assert response.json()["kv_transfer_params"] == kv
# ---------------------------------------------------------------------------
# Resource bounds regression tests
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_derender_chat_bounded_payload_succeeds(client):
"""Normal bounded derender payload succeeds (positive control)."""
gen_req = await _render_chat(client)
synthetic_ids = gen_req["token_ids"][:5]
response = await client.post(
"/v1/chat/completions/derender",
json={
"model": MODEL_NAME,
"generate_response": _make_generate_response(synthetic_ids),
},
)
assert response.status_code == 200
data = response.json()
assert len(data["choices"]) == 1
assert data["choices"][0]["message"]["content"]
@pytest.mark.asyncio
async def test_derender_chat_oversized_token_ids_rejected(client):
"""token_ids longer than max_model_len returns 400."""
response = await client.get("/v1/models")
assert response.status_code == 200
# Use a token_ids list that exceeds any reasonable max_model_len.
# The tiny-random model has max_model_len of 2048.
oversized_ids = [42] * 1_000_000
response = await client.post(
"/v1/chat/completions/derender",
json={
"model": MODEL_NAME,
"generate_response": _make_generate_response(oversized_ids),
},
)
assert response.status_code == 400
assert "max_model_len" in response.json()["error"]["message"]
@pytest.mark.asyncio
async def test_derender_chat_too_many_choices_rejected(client):
"""choices count exceeding VLLM_MAX_N_SEQUENCES returns 400."""
# Default VLLM_MAX_N_SEQUENCES is 16384; use a larger count.
oversized_choices = [
{"index": i, "token_ids": [42], "finish_reason": "stop"} for i in range(20_000)
]
response = await client.post(
"/v1/chat/completions/derender",
json={
"model": MODEL_NAME,
"generate_response": {
"request_id": "test-choices-bound",
"choices": oversized_choices,
},
},
)
assert response.status_code == 400
assert "choices count" in response.json()["error"]["message"]
@pytest.mark.asyncio
async def test_derender_completion_too_many_generate_responses_rejected(client):
"""generate_responses count exceeding limit returns 400."""
oversized_responses = [
{
"request_id": f"gen-{i}",
"choices": [{"index": 0, "token_ids": [42], "finish_reason": "stop"}],
}
for i in range(20_000)
]
response = await client.post(
"/v1/completions/derender",
json={
"model": MODEL_NAME,
"generate_responses": oversized_responses,
},
)
assert response.status_code == 400
assert "generate_responses count" in response.json()["error"]["message"]
@pytest.mark.asyncio
async def test_derender_chat_negative_token_ids_rejected(client):
"""Negative token_ids are rejected at the protocol validation level."""
response = await client.post(
"/v1/chat/completions/derender",
json={
"model": MODEL_NAME,
"generate_response": _make_generate_response([-1, 42, 100]),
},
)
# vLLM's validation_exception_handler converts Pydantic errors to 400
assert response.status_code == 400
@pytest.mark.asyncio
async def test_derender_chat_oversized_logprobs_rejected(client):
"""logprobs.content longer than max_model_len returns 400."""
oversized_logprobs: dict = {
"content": [
{"token": "x", "logprob": -1.0, "bytes": None, "top_logprobs": []}
for _ in range(1_000_000)
]
}
response = await client.post(
"/v1/chat/completions/derender",
json={
"model": MODEL_NAME,
"generate_response": {
"request_id": "test-logprobs-bound",
"choices": [
{
"index": 0,
"token_ids": [42],
"finish_reason": "stop",
"logprobs": oversized_logprobs,
}
],
},
},
)
assert response.status_code == 400
assert "logprobs.content length" in response.json()["error"]["message"]
@pytest.mark.asyncio
async def test_derender_chat_oversized_top_logprobs_rejected(client):
"""top_logprobs count exceeding 20 returns 400."""
oversized_top_logprobs = {
"content": [
{
"token": "x",
"logprob": -1.0,
"bytes": None,
"top_logprobs": [
{"token": f"t{i}", "logprob": -float(i), "bytes": None}
for i in range(25)
],
}
]
}
response = await client.post(
"/v1/chat/completions/derender",
json={
"model": MODEL_NAME,
"generate_response": {
"request_id": "test-top-logprobs-bound",
"choices": [
{
"index": 0,
"token_ids": [42],
"finish_reason": "stop",
"logprobs": oversized_top_logprobs,
}
],
},
},
)
assert response.status_code == 400
assert "top_logprobs count" in response.json()["error"]["message"]
@pytest.mark.asyncio
async def test_derender_completion_oversized_token_ids_rejected(client):
"""Completion endpoint also rejects oversized token_ids."""
oversized_ids = [42] * 1_000_000
response = await client.post(
"/v1/completions/derender",
json={
"model": MODEL_NAME,
"generate_responses": [
{
"request_id": "gen-0",
"choices": [
{
"index": 0,
"token_ids": oversized_ids,
"finish_reason": "stop",
}
],
}
],
},
)
assert response.status_code == 400
assert "max_model_len" in response.json()["error"]["message"]
# ---------------------------------------------------------------------------
# E2E: render -> derender roundtrip with parser (reasoning + tool calls)
# ---------------------------------------------------------------------------
+71 -3
View File
@@ -3,6 +3,7 @@
import time
from typing import cast
import vllm.envs as envs
from vllm.entrypoints.openai.chat_completion.protocol import ChatCompletionResponse
from vllm.entrypoints.openai.completion.protocol import CompletionResponse
from vllm.entrypoints.openai.engine.protocol import (
@@ -28,6 +29,7 @@ from ..token_in_token_out.mm_serde import encode_mm_kwargs_item
from ..token_in_token_out.protocol import (
DerenderChatRequest,
DerenderCompletionRequest,
GenerateResponse,
MultiModalFeatures,
PlaceholderRangeInfo,
)
@@ -51,6 +53,64 @@ class ServingDerender(BaseServing):
self.online_derenderer = online_derenderer
def _validate_derender_bounds(
self,
generate_responses: list[GenerateResponse],
) -> ErrorResponse | None:
"""Reject derender payloads that exceed resource bounds.
Runs before any tokenizer.decode() or parser invocation to prevent
CPU/memory exhaustion from oversized caller-supplied token structures.
"""
max_n = envs.VLLM_MAX_N_SEQUENCES
max_model_len = self.model_config.max_model_len
if len(generate_responses) > max_n:
return self.create_error_response(
f"generate_responses count ({len(generate_responses)}) "
f"exceeds server maximum ({max_n}). "
f"Set VLLM_MAX_N_SEQUENCES to increase this limit."
)
for gen in generate_responses:
if len(gen.choices) > max_n:
return self.create_error_response(
f"choices count ({len(gen.choices)}) in response "
f"'{gen.request_id}' exceeds server maximum ({max_n})."
)
for choice in gen.choices:
if choice.token_ids and len(choice.token_ids) > max_model_len:
return self.create_error_response(
f"token_ids length ({len(choice.token_ids)}) in "
f"choice {choice.index} exceeds "
f"max_model_len ({max_model_len})."
)
if choice.logprobs and choice.logprobs.content:
if len(choice.logprobs.content) > max_model_len:
return self.create_error_response(
f"logprobs.content length "
f"({len(choice.logprobs.content)}) in "
f"choice {choice.index} exceeds "
f"max_model_len ({max_model_len})."
)
for entry in choice.logprobs.content:
if entry.top_logprobs and len(entry.top_logprobs) > 20:
return self.create_error_response(
f"top_logprobs count "
f"({len(entry.top_logprobs)}) in "
f"choice {choice.index} exceeds maximum (20)."
)
if gen.prompt_logprobs and len(gen.prompt_logprobs) > max_model_len:
return self.create_error_response(
f"prompt_logprobs length ({len(gen.prompt_logprobs)}) "
f"in response '{gen.request_id}' exceeds "
f"max_model_len ({max_model_len})."
)
return None
async def derender_chat_response(
self,
request: DerenderChatRequest,
@@ -68,6 +128,10 @@ class ServingDerender(BaseServing):
if error_check_ret is not None:
return error_check_ret
bounds_error = self._validate_derender_bounds([request.generate_response])
if bounds_error is not None:
return bounds_error
try:
choices = await self.online_derenderer.derender_chat(
request.generate_response, request.chat_request
@@ -117,6 +181,13 @@ class ServingDerender(BaseServing):
if error_check_ret is not None:
return error_check_ret
if not request.generate_responses:
return self.create_error_response("generate_responses must not be empty")
bounds_error = self._validate_derender_bounds(request.generate_responses)
if bounds_error is not None:
return bounds_error
(
choices,
total_prompt_tokens,
@@ -125,9 +196,6 @@ class ServingDerender(BaseServing):
request.generate_responses, request.prompt_tokens
)
if not request.generate_responses:
return self.create_error_response("generate_responses must not be empty")
first = request.generate_responses[0]
kv_params = first.kv_transfer_params
if any(
@@ -190,6 +190,13 @@ class GenerateResponseChoice(BaseModel):
# or (b) ``enable_return_routed_experts`` is off server-side.
routed_experts: str | None = None
@field_validator("token_ids")
@classmethod
def validate_token_ids(cls, v: list[int] | None) -> list[int] | None:
if v is not None and any(t < 0 for t in v):
raise ValueError("token_ids must not contain negative values")
return v
class GenerateResponseStreamChoice(BaseModel):
index: int