diff --git a/docs/features/disagg_prefill.md b/docs/features/disagg_prefill.md index 4c5265449c97..e6a5797f268a 100644 --- a/docs/features/disagg_prefill.md +++ b/docs/features/disagg_prefill.md @@ -77,6 +77,8 @@ decode = client.chat.completions.create( ) ``` +If `messages` has non-text content or `echo` is set, the ids are ignored and `messages` is rendered instead, so it must match the prefill request. Otherwise `kv_transfer_params["prompt_token_ids"]` must be a non-empty list of non-negative integers, or the request fails with HTTP 400, as it always does on `/v1/chat/completions/batch`. + ## Development We implement disaggregated prefilling by running 2 vLLM instances. One for prefill (we call it prefill instance) and one for decode (we call it decode instance), and then use a connector to transfer the prefill KV caches and results from prefill instance to decode instance. diff --git a/tests/entrypoints/openai/chat_completion/test_batched_chat_completions.py b/tests/entrypoints/openai/chat_completion/test_batched_chat_completions.py index f1de27bf756d..a6d78c826cad 100644 --- a/tests/entrypoints/openai/chat_completion/test_batched_chat_completions.py +++ b/tests/entrypoints/openai/chat_completion/test_batched_chat_completions.py @@ -26,6 +26,7 @@ BatchChatCompletionRequest, ) from vllm.entrypoints.openai.models.serving import OpenAIServingModels +from vllm.exceptions import VLLMValidationError from vllm.outputs import CompletionOutput, RequestOutput from vllm.renderers.online_renderer import OnlineRenderer from vllm.v1.engine.async_llm import AsyncLLM @@ -385,3 +386,19 @@ async def test_batched_harmony_response_format_uses_structural_tag() -> None: assert len(calls) == 2 for call in calls: assert call.args[1].structured_outputs.structural_tag is not None + + +@pytest.mark.skip_global_cleanup +def test_batch_rejects_kv_transfer_prompt_token_ids(): + """One pre-tokenized prompt cannot stand in for every conversation.""" + with pytest.raises( + VLLMValidationError, match=r"parameter=kv_transfer_params\.prompt_token_ids" + ): + BatchChatCompletionRequest( + model="test-model", + messages=[ + [{"role": "user", "content": "first"}], + [{"role": "user", "content": "second"}], + ], + kv_transfer_params={"prompt_token_ids": [10, 20, 30]}, + ) diff --git a/tests/entrypoints/openai/chat_completion/test_serving_chat.py b/tests/entrypoints/openai/chat_completion/test_serving_chat.py index 6022f5ce7733..241920d1e3df 100644 --- a/tests/entrypoints/openai/chat_completion/test_serving_chat.py +++ b/tests/entrypoints/openai/chat_completion/test_serving_chat.py @@ -595,6 +595,15 @@ def _build_online_renderer( ) +def _build_mock_engine() -> MagicMock: + mock_engine = MagicMock(spec=AsyncLLM) + mock_engine.errored = False + mock_engine.model_config = MockModelConfig() + mock_engine.input_processor = MagicMock() + mock_engine.renderer = _build_renderer(mock_engine.model_config) + return mock_engine + + def _build_serving_chat( engine: AsyncLLM, *, @@ -2817,3 +2826,132 @@ def test_make_request_with_harmony_reuses_kv_transfer_prompt_token_ids(): assert engine_input["prompt_token_ids"] == [10, 20, 30] # The reuse key is consumed and other kv_transfer_params are preserved. assert request.kv_transfer_params == {"do_remote_prefill": True} + + +@pytest.mark.parametrize("ids", [[-1], [1.5]]) +def test_make_request_with_harmony_rejects_invalid_kv_transfer_prompt_token_ids(ids): + engine = MockEngine() + engine.model_config.hf_config = MockHFConfig(model_type="gpt_oss") + models = OpenAIServingModels(engine, BASE_MODEL_PATHS) + online_renderer = _build_online_renderer(engine, models.registry) + + request = ChatCompletionRequest( + model=MODEL_NAME, + messages=[{"role": "user", "content": "hi"}], + kv_transfer_params={"prompt_token_ids": ids}, + ) + with pytest.raises(VLLMValidationError, match="non-negative integers"): + online_renderer._make_request_with_harmony(request) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("ids", [[], [-1], [1.5], [1.0], ["1"], [True], "abc", 5]) +async def test_chat_kv_transfer_prompt_token_ids_rejects_invalid_ids(ids): + """``kv_transfer_params`` is untyped, so the ids are checked on reuse.""" + serving_chat = _build_serving_chat(_build_mock_engine()) + + request = ChatCompletionRequest( + model=MODEL_NAME, + messages=[{"role": "user", "content": "hi"}], + kv_transfer_params={"prompt_token_ids": ids}, + ) + with pytest.raises(VLLMValidationError, match="non-negative integers"): + await serving_chat.render_chat_request(request) + + +@pytest.mark.asyncio +async def test_chat_kv_transfer_prompt_token_ids_rejection_notifies_kv_connector(): + """A render-time 400 on a decode request notifies the KV connector.""" + mock_engine = _build_mock_engine() + mock_engine.notify_kv_transfer_request_rejected = AsyncMock() + serving_chat = _build_serving_chat(mock_engine) + serving_chat.has_kv_connector = True + + request = ChatCompletionRequest( + model=MODEL_NAME, + messages=[{"role": "user", "content": "hi"}], + kv_transfer_params={"prompt_token_ids": [1.5], "do_remote_prefill": True}, + ) + with pytest.raises(VLLMValidationError, match="non-negative integers"): + await serving_chat.create_chat_completion(request) + + mock_engine.notify_kv_transfer_request_rejected.assert_awaited_once_with( + request.request_id, {"do_remote_prefill": True}, data_parallel_rank=None + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("ids", [[10, 20, 30], [1.5]]) +async def test_chat_kv_transfer_prompt_token_ids_ignored_with_echo(ids): + """``echo`` needs the rendered conversation, so the ids are not used.""" + serving_chat = _build_serving_chat(_build_mock_engine()) + + request = ChatCompletionRequest( + model=MODEL_NAME, + messages=[{"role": "user", "content": "hi"}], + kv_transfer_params={"prompt_token_ids": ids, "do_remote_prefill": True}, + echo=True, + ) + result = await serving_chat.render_chat_request(request) + assert not isinstance(result, ErrorResponse) + + conversation, engine_inputs = result + assert [msg["role"] for msg in conversation] == ["user"] + assert engine_inputs[0]["prompt_token_ids"] != ids + assert request.kv_transfer_params == {"do_remote_prefill": True} + + +@pytest.mark.parametrize( + "content", + [ + [ + {"type": "text", "text": "what is in this image?"}, + {"type": "image_url", "image_url": {"url": "https://example.com/a.png"}}, + ], + # A part written without a "type" is identified by its media key. + [{"image_url": "https://example.com/a.png"}], + # A media key counts even when the part claims to be text, which is + # how chat_utils reads a part carrying a "uuid". + [ + { + "type": "text", + "text": "look", + "uuid": "u1", + "image_url": "https://example.com/a.png", + } + ], + [{"type": "hologram", "hologram": "https://a/b.holo"}], + [{"type": ["image_url"]}], + ], +) +def test_chat_kv_transfer_prompt_token_ids_ignored_with_non_text_content(content): + """The ids are dropped so that ``messages`` is rendered with its media.""" + request = ChatCompletionRequest( + model=MODEL_NAME, + messages=[{"role": "user", "content": content}], + kv_transfer_params={ + "prompt_token_ids": [10, 20, 30], + "do_remote_prefill": True, + }, + ) + assert request.kv_transfer_params == {"do_remote_prefill": True} + + +def test_chat_kv_transfer_prompt_token_ids_allows_text_only_parts(): + """Text-bearing part types are not multimodal, so the ids are kept.""" + request = ChatCompletionRequest( + model=MODEL_NAME, + messages=[ + {"role": "user", "content": [{"type": "text", "text": "hi"}]}, + { + "role": "assistant", + "content": [ + {"type": "thinking", "thinking": "let me think"}, + {"type": "refusal", "refusal": "no"}, + ], + }, + {"role": "user", "content": "plain string"}, + ], + kv_transfer_params={"prompt_token_ids": [10, 20, 30]}, + ) + assert request.kv_transfer_params == {"prompt_token_ids": [10, 20, 30]} diff --git a/vllm/entrypoints/chat_utils.py b/vllm/entrypoints/chat_utils.py index f739495a8a1f..a419ffca81b7 100644 --- a/vllm/entrypoints/chat_utils.py +++ b/vllm/entrypoints/chat_utils.py @@ -1810,6 +1810,11 @@ def _parse_chat_message_content_mm_part( "refusal", ) +# Content part types parsed as text rather than multimodal data. +TEXT_PART_TYPES = frozenset( + {"text", "input_text", "output_text", "refusal", "thinking"} +) + def _parse_chat_message_content_parts( role: str, @@ -1904,7 +1909,7 @@ def _parse_chat_message_content_part( ) return None - if part_type in ("text", "input_text", "output_text", "refusal", "thinking"): + if part_type in TEXT_PART_TYPES: str_content = cast(str, content) _reject_reserved_placeholder_in_text(str_content, mm_parser.model_config) if wrap_dicts: diff --git a/vllm/entrypoints/openai/chat_completion/protocol.py b/vllm/entrypoints/openai/chat_completion/protocol.py index 5c1bc594d70e..8ad7534d6dfe 100644 --- a/vllm/entrypoints/openai/chat_completion/protocol.py +++ b/vllm/entrypoints/openai/chat_completion/protocol.py @@ -20,6 +20,8 @@ from vllm.config import ModelConfig from vllm.entrypoints.chat_utils import ( + MM_PARSER_MAP, + TEXT_PART_TYPES, ChatCompletionMessageParam, ChatTemplateContentFormatOption, ) @@ -58,6 +60,11 @@ _INT64_MIN = -(2**63) _INT64_MAX = 2**63 - 1 +# Content part types that carry no multimodal data. +_TEXT_CONTENT_PART_TYPES = TEXT_PART_TYPES | {"tool_reference"} +# Keys that mark a content part as multimodal, whatever its ``type``. +_MEDIA_CONTENT_PART_KEYS = frozenset(MM_PARSER_MAP) - _TEXT_CONTENT_PART_TYPES + class ChatMessage(OpenAIBaseModel): role: str @@ -976,6 +983,39 @@ def check_generation_prompt(cls, data): ) return data + @model_validator(mode="before") + @classmethod + def drop_prompt_token_ids_with_media(cls, data): + # The forwarded ids would drop media in ``messages``, so ignore them. Runs + # before validation, which can turn content into a one-shot iterator. + if not isinstance(data, dict): + return data + kv_transfer_params = data.get("kv_transfer_params") + messages = data.get("messages") + if ( + not isinstance(kv_transfer_params, dict) + or kv_transfer_params.get("prompt_token_ids") is None + or not isinstance(messages, list) + ): + return data + for msg in messages: + content = msg.get("content") if isinstance(msg, dict) else None + if not isinstance(content, list): + continue + for part in content: + if isinstance(part, dict) and ( + any(key in part for key in _MEDIA_CONTENT_PART_KEYS) + or not isinstance(part_type := part.get("type", "text"), str) + or part_type not in _TEXT_CONTENT_PART_TYPES + ): + logger.debug( + "Ignoring kv_transfer_params['prompt_token_ids']: " + "messages have non-text content and are rendered instead." + ) + kv_transfer_params.pop("prompt_token_ids") + return data + return data + @model_validator(mode="before") @classmethod def check_system_message_content_type(cls, data): @@ -1117,6 +1157,17 @@ def check_batch_mode(cls, data: Any) -> Any: "when using `logprob_token_ids`, `logprobs` must be set to true.", parameter="logprob_token_ids", ) + kv_transfer_params = data.get("kv_transfer_params") + if ( + isinstance(kv_transfer_params, dict) + and kv_transfer_params.get("prompt_token_ids") is not None + ): + raise VLLMValidationError( + "Batch chat completions do not support " + "`kv_transfer_params['prompt_token_ids']`: one pre-tokenized " + "prompt cannot serve several conversations.", + parameter="kv_transfer_params.prompt_token_ids", + ) response_format = data.get("response_format") if response_format is not None: rf_type = ( diff --git a/vllm/renderers/online_renderer.py b/vllm/renderers/online_renderer.py index be0056a3506d..5108adfe81f0 100644 --- a/vllm/renderers/online_renderer.py +++ b/vllm/renderers/online_renderer.py @@ -107,11 +107,35 @@ def _reused_prompt_token_ids(request: Any) -> list[int] | None: Disaggregated serving carries the prefill stage's ids in ``kv_transfer_params`` so the decode stage can skip re-tokenizing. Removing the key keeps the id list out of the engine's sampling metadata. + + Returns None without checking the ids when ``echo`` is set, since echo + needs ``messages`` to be rendered. Otherwise raises VLLMValidationError if + the ids are malformed. """ kv = getattr(request, "kv_transfer_params", None) if not isinstance(kv, dict): return None - return kv.pop("prompt_token_ids", None) or None + ids = kv.pop("prompt_token_ids", None) + if ids is None: + return None + if getattr(request, "echo", False): + logger.debug( + "Ignoring kv_transfer_params['prompt_token_ids']: " + "echo is set, so messages are rendered instead." + ) + return None + # bool is an int subclass, hence the exact type check. + if ( + not isinstance(ids, list) + or not ids + or any(type(x) is not int or x < 0 for x in ids) + ): + raise VLLMValidationError( + "`kv_transfer_params['prompt_token_ids']` must be a non-empty list " + "of non-negative integers.", + parameter="kv_transfer_params.prompt_token_ids", + ) + return ids class OnlineRenderer: