diff --git a/tests/entrypoints/scale_out/derender/test_derender.py b/tests/entrypoints/scale_out/derender/test_derender.py index e452b7367a2..880d9396264 100644 --- a/tests/entrypoints/scale_out/derender/test_derender.py +++ b/tests/entrypoints/scale_out/derender/test_derender.py @@ -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) # --------------------------------------------------------------------------- diff --git a/vllm/entrypoints/scale_out/derender/serving.py b/vllm/entrypoints/scale_out/derender/serving.py index e125007a549..f46ca7f4a57 100644 --- a/vllm/entrypoints/scale_out/derender/serving.py +++ b/vllm/entrypoints/scale_out/derender/serving.py @@ -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( diff --git a/vllm/entrypoints/scale_out/token_in_token_out/protocol.py b/vllm/entrypoints/scale_out/token_in_token_out/protocol.py index 48a1c4722dd..233ebf070c5 100644 --- a/vllm/entrypoints/scale_out/token_in_token_out/protocol.py +++ b/vllm/entrypoints/scale_out/token_in_token_out/protocol.py @@ -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