forked from Karylab-cklius/vllm
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:
co-authored by
mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
parent
e040899a00
commit
8e61b646e2
@@ -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)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user