Skip to content
2 changes: 2 additions & 0 deletions docs/features/disagg_prefill.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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]},
)
138 changes: 138 additions & 0 deletions tests/entrypoints/openai/chat_completion/test_serving_chat.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
*,
Expand Down Expand Up @@ -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]}
7 changes: 6 additions & 1 deletion vllm/entrypoints/chat_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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:
Expand Down
51 changes: 51 additions & 0 deletions vllm/entrypoints/openai/chat_completion/protocol.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,8 @@

from vllm.config import ModelConfig
from vllm.entrypoints.chat_utils import (
MM_PARSER_MAP,
TEXT_PART_TYPES,
ChatCompletionMessageParam,
ChatTemplateContentFormatOption,
)
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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")
Comment thread
shijie-lyu marked this conversation as resolved.
return data
return data

@model_validator(mode="before")
@classmethod
def check_system_message_content_type(cls, data):
Expand Down Expand Up @@ -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 = (
Expand Down
26 changes: 25 additions & 1 deletion vllm/renderers/online_renderer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
):
Comment thread
shijie-lyu marked this conversation as resolved.
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:
Expand Down
Loading