[Frontend] Add detokenization streaming derender for disaggregated serving (#47301)

Signed-off-by: Martin Hickey <[email protected]>
This commit is contained in:
Martin Hickey
2026-07-29 18:01:20 -07:00
committed by GitHub
parent 451227cb3f
commit 0a31372e5f
6 changed files with 1377 additions and 32 deletions
@@ -0,0 +1,765 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Unit tests for streaming derender.
Tests are split into two layers:
1. Unit tests (no server): covers ``_detokenize_delta`` correctness
(chunked == one-shot) and ``derender_completion_stream`` /
``derender_chat_stream`` logic via a real tokenizer on a tiny model.
2. Integration tests (require a running render server): covers the full
HTTP round-trip through the streaming endpoint. Marked with
``@pytest.mark.asyncio`` and gated by the ``server`` / ``client``
fixtures from the sibling ``test_derender.py``.
"""
import pytest
import pytest_asyncio
from vllm.entrypoints.scale_out.token_in_token_out.protocol import (
DerenderStreamState,
GenerateResponseStreamChoice,
GenerateStreamResponse,
)
MODEL_NAME = "hmellor/tiny-random-LlamaForCausalLM"
# ---------------------------------------------------------------------------
# Helpers shared across tests
# ---------------------------------------------------------------------------
def _make_stream_chunk(
token_ids: list[int],
index: int = 0,
finish_reason: str | None = None,
request_id: str = "test-req",
usage: dict | None = None,
) -> GenerateStreamResponse:
"""Build a GenerateStreamResponse SSE chunk."""
from vllm.entrypoints.openai.engine.protocol import UsageInfo
return GenerateStreamResponse(
request_id=request_id,
choices=[
GenerateResponseStreamChoice(
index=index,
token_ids=token_ids,
finish_reason=finish_reason,
)
],
usage=UsageInfo(**usage) if usage else None,
)
def _make_usage_chunk(
completion_tokens: int,
prompt_tokens: int = 0,
request_id: str = "test-req",
) -> GenerateStreamResponse:
"""Build a usage only final SSE chunk (empty choices)."""
from vllm.entrypoints.openai.engine.protocol import UsageInfo
return GenerateStreamResponse(
request_id=request_id,
choices=[],
usage=UsageInfo(
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
total_tokens=prompt_tokens + completion_tokens,
),
)
# ---------------------------------------------------------------------------
# Unit tests — no running server
# ---------------------------------------------------------------------------
@pytest.fixture(scope="module")
def tokenizer():
"""Load the tiny tokenizer used across unit tests."""
from vllm.tokenizers import get_tokenizer
return get_tokenizer(MODEL_NAME)
@pytest.fixture(scope="module")
def derenderer(tokenizer):
"""Construct a minimal OnlineDerenderer backed by a stub renderer."""
from unittest.mock import MagicMock
from vllm.renderers.online_derenderer import OnlineDerenderer
renderer = MagicMock()
renderer.get_tokenizer.return_value = tokenizer
model_config = MagicMock()
model_config.hf_config.model_type = "llama"
model_config.model = MODEL_NAME
return OnlineDerenderer(
model_config=model_config,
renderer=renderer,
request_logger=None,
chat_template=None,
chat_template_content_format="string",
trust_request_chat_template=False,
enable_auto_tools=False,
tool_parser=None,
reasoning_parser=None,
)
class TestDetokenizeDelta:
"""_detokenize_delta: chunked decode must equal one shot decode."""
def _one_shot(self, tokenizer, token_ids: list[int]) -> str:
return tokenizer.decode(token_ids, skip_special_tokens=True)
def _chunked(self, derenderer, tokenizer, chunks: list[list[int]]) -> str:
state = DerenderStreamState()
parts: list[str] = []
for delta in chunks:
text, state = derenderer._detokenize_delta(
tokenizer, delta, state, skip_special_tokens=True
)
parts.append(text)
return "".join(parts)
def test_single_chunk(self, derenderer, tokenizer):
"""All tokens in one chunk == one shot decode."""
token_ids = tokenizer.encode("Hello world")[:8]
assert self._chunked(derenderer, tokenizer, [token_ids]) == self._one_shot(
tokenizer, token_ids
)
def test_two_equal_chunks(self, derenderer, tokenizer):
"""Split in half and reassemble == one shot."""
token_ids = tokenizer.encode("Hello world from streaming derender")[:12]
mid = len(token_ids) // 2
chunks = [token_ids[:mid], token_ids[mid:]]
assert self._chunked(derenderer, tokenizer, chunks) == self._one_shot(
tokenizer, token_ids
)
def test_single_token_per_chunk(self, derenderer, tokenizer):
"""One token per chunk (most granular streaming) == one shot."""
token_ids = tokenizer.encode("incremental detokenization test")[:10]
chunks = [[t] for t in token_ids]
assert self._chunked(derenderer, tokenizer, chunks) == self._one_shot(
tokenizer, token_ids
)
def test_empty_delta_passthrough(self, derenderer, tokenizer):
"""An empty delta (usage only chunk) emits empty string and preserves state."""
token_ids = tokenizer.encode("Hello")[:4]
_, state = derenderer._detokenize_delta(
tokenizer, token_ids, DerenderStreamState(), skip_special_tokens=True
)
text, new_state = derenderer._detokenize_delta(
tokenizer, [], state, skip_special_tokens=True
)
assert text == ""
assert new_state.prev_tokens == state.prev_tokens
assert new_state.prefix_offset == state.prefix_offset
assert new_state.read_offset == state.read_offset
def test_multibyte_char_split_across_chunks(self, derenderer, tokenizer):
"""A CJK/emoji char straddling chunk boundaries == one shot.
Regression test for held back trailing incomplete UTF-8 byte
sequences being dropped when the rebuild window marks them as
already read (see #46159).
"""
token_ids = tokenizer.encode("Hello ✅ world 日本語")[:16]
chunks = [[t] for t in token_ids]
assert self._chunked(derenderer, tokenizer, chunks) == self._one_shot(
tokenizer, token_ids
)
def test_state_carries_across_calls(self, derenderer, tokenizer):
"""Decode state threads across calls. Text still matches one shot."""
t1 = tokenizer.encode("Hello")[:2]
t2 = tokenizer.encode(" world")[:2]
state = DerenderStreamState()
text1, state = derenderer._detokenize_delta(tokenizer, t1, state)
text2, state = derenderer._detokenize_delta(tokenizer, t2, state)
assert text1 + text2 == self._one_shot(tokenizer, t1 + t2)
# Offsets are rebased to the carried tail each chunk
assert state.prefix_offset == 0
def test_state_window_stays_bounded(self, derenderer, tokenizer):
"""prev_tokens must not grow with the number of chunks (bounded transport).
Guards that the carried decode window is a small constant
tail, so cumulative ``stream_state`` transport is O(n) and not O(n^2).
"""
token_ids = tokenizer.encode(
"a reasonably long ascii stream of tokens used to exercise the "
"window bound across many single token chunks so the carried "
"state cannot grow linearly with the generation length"
)
assert len(token_ids) > 32
state = DerenderStreamState()
max_window = 0
for tok in token_ids:
_, state = derenderer._detokenize_delta(tokenizer, [tok], state)
max_window = max(max_window, len(state.prev_tokens))
# Bounded by a small constant independent of len(token_ids)
assert max_window <= 32
def test_n_independent_streams_same_result(self, derenderer, tokenizer):
"""N parallel streams with the same token sequence give the same text."""
token_ids = tokenizer.encode("parallel streams")[:8]
mid = len(token_ids) // 2
results = []
for _ in range(3):
state = DerenderStreamState()
text, state = derenderer._detokenize_delta(
tokenizer, token_ids[:mid], state
)
text2, _ = derenderer._detokenize_delta(tokenizer, token_ids[mid:], state)
results.append(text + text2)
assert len(set(results)) == 1, "All independent streams must produce same text"
assert results[0] == self._one_shot(tokenizer, token_ids)
class TestDerenderCompletionStream:
"""derender_completion_stream: streaming output parity with one shot."""
@pytest.mark.asyncio
async def test_chunked_equals_oneshot(self, derenderer, tokenizer):
"""Sum of streaming text chunks == one shot tokenizer.decode."""
token_ids = tokenizer.encode("streaming completion test")[:10]
mid = len(token_ids) // 2
state = DerenderStreamState()
chunk1, state = await derenderer.derender_completion_stream(
model=MODEL_NAME,
generate_chunk=_make_stream_chunk(token_ids[:mid]),
state=state,
)
chunk2, _ = await derenderer.derender_completion_stream(
model=MODEL_NAME,
generate_chunk=_make_stream_chunk(token_ids[mid:], finish_reason="stop"),
state=state,
)
streamed_text = chunk1.choices[0].text + chunk2.choices[0].text
one_shot = tokenizer.decode(token_ids, skip_special_tokens=True)
assert streamed_text == one_shot
@pytest.mark.asyncio
async def test_usage_chunk_passthrough(self, derenderer, tokenizer):
"""Usage only final chunk (empty choices) is passed through correctly."""
usage_chunk = _make_usage_chunk(completion_tokens=10, prompt_tokens=5)
chunk, state = await derenderer.derender_completion_stream(
model=MODEL_NAME,
generate_chunk=usage_chunk,
)
assert chunk.choices == []
assert chunk.usage is not None
assert chunk.usage.completion_tokens == 10
assert chunk.usage.prompt_tokens == 5
@pytest.mark.asyncio
async def test_prompt_tokens_in_usage(self, derenderer, tokenizer):
"""prompt_tokens is correctly forwarded into usage on a usage chunk."""
token_ids = tokenizer.encode("hello")[:3]
usage_chunk = _make_usage_chunk(
completion_tokens=len(token_ids), prompt_tokens=7
)
chunk, _ = await derenderer.derender_completion_stream(
model=MODEL_NAME,
generate_chunk=usage_chunk,
prompt_tokens=7,
)
assert chunk.usage is not None
assert chunk.usage.prompt_tokens == 7
@pytest.mark.asyncio
async def test_none_state_initialises_correctly(self, derenderer, tokenizer):
"""Passing state=None (first call) initialises an empty DerenderStreamState."""
token_ids = tokenizer.encode("hello")[:4]
chunk, state = await derenderer.derender_completion_stream(
model=MODEL_NAME,
generate_chunk=_make_stream_chunk(token_ids),
state=None,
)
assert isinstance(state, DerenderStreamState)
assert chunk.choices[0].text == tokenizer.decode(
token_ids, skip_special_tokens=True
)
@pytest.mark.asyncio
async def test_skip_special_tokens_threaded(self, derenderer, tokenizer):
"""completion_request.skip_special_tokens is honored (not hardcoded True)."""
from vllm.entrypoints.openai.completion.protocol import CompletionRequest
eos = tokenizer.eos_token_id
if eos is None:
pytest.skip("tokenizer has no eos token to exercise special stripping")
token_ids = tokenizer.encode("hi")[:2] + [eos]
async def _text(skip: bool) -> str:
req = CompletionRequest(
model=MODEL_NAME, prompt="x", skip_special_tokens=skip
)
chunk, _ = await derenderer.derender_completion_stream(
model=MODEL_NAME,
generate_chunk=_make_stream_chunk(token_ids),
completion_request=req,
)
return chunk.choices[0].text
# skip=False must retain the special token; skip=True must strip it.
assert await _text(False) != await _text(True)
@pytest.mark.asyncio
async def test_finish_reason_forwarded(self, derenderer, tokenizer):
"""finish_reason from the generate chunk reaches the derendered choice."""
token_ids = tokenizer.encode("done")[:2]
chunk, _ = await derenderer.derender_completion_stream(
model=MODEL_NAME,
generate_chunk=_make_stream_chunk(token_ids, finish_reason="length"),
)
assert chunk.choices[0].finish_reason == "length"
class TestDerenderChatStream:
"""derender_chat_stream: plain detok branch (no parser)."""
@pytest.mark.asyncio
async def test_role_on_first_chunk_only(self, derenderer, tokenizer):
"""role='assistant' appears in the first chunk, not subsequent ones."""
token_ids = tokenizer.encode("hello world")[:6]
mid = len(token_ids) // 2
state = DerenderStreamState()
chunk1, state = await derenderer.derender_chat_stream(
model=MODEL_NAME,
generate_chunk=_make_stream_chunk(token_ids[:mid]),
state=state,
)
chunk2, _ = await derenderer.derender_chat_stream(
model=MODEL_NAME,
generate_chunk=_make_stream_chunk(token_ids[mid:], finish_reason="stop"),
state=state,
)
assert chunk1.choices[0].delta.role == "assistant"
assert chunk2.choices[0].delta.role is None
@pytest.mark.asyncio
async def test_chunked_equals_oneshot(self, derenderer, tokenizer):
"""Sum of streaming content deltas == one shot decode."""
token_ids = tokenizer.encode("streaming chat derender text")[:10]
mid = len(token_ids) // 2
state = DerenderStreamState()
chunk1, state = await derenderer.derender_chat_stream(
model=MODEL_NAME,
generate_chunk=_make_stream_chunk(token_ids[:mid]),
state=state,
)
chunk2, _ = await derenderer.derender_chat_stream(
model=MODEL_NAME,
generate_chunk=_make_stream_chunk(token_ids[mid:]),
state=state,
)
streamed = (chunk1.choices[0].delta.content or "") + (
chunk2.choices[0].delta.content or ""
)
one_shot = tokenizer.decode(token_ids, skip_special_tokens=True)
assert streamed == one_shot
@pytest.mark.asyncio
@pytest.mark.parametrize("with_chat_request", [True, False])
async def test_parser_raises_not_implemented(self, tokenizer, with_chat_request):
"""Stream chat with a parser active must fail closed (NotImplementedError).
Checks that a parser configured model must 501 even when ``chat_request`` is
omitted, otherwise reasoning/tool markup would leak into ``delta.content`
via the plain detok fallback.
"""
from unittest.mock import MagicMock
from vllm.renderers.online_derenderer import OnlineDerenderer
renderer = MagicMock()
renderer.get_tokenizer.return_value = tokenizer
model_config = MagicMock()
model_config.hf_config.model_type = "llama"
model_config.model = MODEL_NAME
# Construct a derenderer WITH a tool/reasoning parser active.
# ParserManager.get_parser returns None when no parser is named, so
# we inject a mock parser class directly.
dr = OnlineDerenderer(
model_config=model_config,
renderer=renderer,
request_logger=None,
chat_template=None,
chat_template_content_format="string",
)
dr.parser = MagicMock() # simulate a parser being active
chat_request = MagicMock() if with_chat_request else None
token_ids = tokenizer.encode("hello")[:3]
with pytest.raises(NotImplementedError):
await dr.derender_chat_stream(
model=MODEL_NAME,
generate_chunk=_make_stream_chunk(token_ids),
state=None,
chat_request=chat_request,
)
class TestDerenderStreamStateValidation:
"""DerenderStreamState rejects malformed caller supplied offsets/lengths."""
def test_negative_prefix_offset_rejected(self):
from pydantic import ValidationError
with pytest.raises(ValidationError):
DerenderStreamState(prefix_offset=-1)
def test_negative_read_offset_rejected(self):
from pydantic import ValidationError
with pytest.raises(ValidationError):
DerenderStreamState(read_offset=-1)
def test_prev_tokens_over_cap_rejected(self):
from pydantic import ValidationError
with pytest.raises(ValidationError):
DerenderStreamState(prev_tokens=["a"] * 1025)
def test_prev_tokens_at_cap_accepted(self):
state = DerenderStreamState(prev_tokens=["a"] * 1024)
assert len(state.prev_tokens) == 1024
class TestServingDerenderStreamErrorHandling:
"""Malformed stream_state must surface as 400 and not an unhandled 500."""
def _make_serving(self, side_effect: Exception):
from unittest.mock import AsyncMock, MagicMock
from vllm.entrypoints.scale_out.derender.serving import ServingDerender
models = MagicMock()
models.is_base_model.return_value = True
models.model_config = MagicMock()
online_derenderer = MagicMock()
online_derenderer.derender_completion_stream = AsyncMock(
side_effect=side_effect
)
online_derenderer.derender_chat_stream = AsyncMock(side_effect=side_effect)
return ServingDerender(models=models, online_derenderer=online_derenderer)
@pytest.mark.asyncio
@pytest.mark.parametrize("exc", [KeyError("bad byte"), IndexError("oob")])
async def test_completion_stream_bad_state_returns_400(self, exc):
from vllm.entrypoints.openai.engine.protocol import ErrorResponse
from vllm.entrypoints.scale_out.token_in_token_out.protocol import (
DerenderCompletionStreamRequest,
)
serving = self._make_serving(exc)
request = DerenderCompletionStreamRequest(
stream=True,
model=MODEL_NAME,
generate_chunk=_make_stream_chunk([1, 2]),
stream_state=DerenderStreamState(),
)
result = await serving.derender_completion_stream_response(request)
assert isinstance(result, ErrorResponse)
assert result.error.code == 400
@pytest.mark.asyncio
@pytest.mark.parametrize("exc", [KeyError("bad byte"), IndexError("oob")])
async def test_chat_stream_bad_state_returns_400(self, exc):
from vllm.entrypoints.openai.engine.protocol import ErrorResponse
from vllm.entrypoints.scale_out.token_in_token_out.protocol import (
DerenderChatStreamRequest,
)
serving = self._make_serving(exc)
request = DerenderChatStreamRequest(
stream=True,
model=MODEL_NAME,
generate_chunk=_make_stream_chunk([1, 2]),
stream_state=DerenderStreamState(),
)
result = await serving.derender_chat_stream_response(request)
assert isinstance(result, ErrorResponse)
assert result.error.code == 400
# ---------------------------------------------------------------------------
# Integration tests — require a live render server
# ---------------------------------------------------------------------------
@pytest.fixture(scope="module")
def server():
from tests.utils import RemoteLaunchRenderServer
with RemoteLaunchRenderServer(MODEL_NAME, []) as remote_server:
yield remote_server
@pytest_asyncio.fixture
async def client(server):
import httpx
async with httpx.AsyncClient(
base_url=server.url_for(""), timeout=30.0
) as http_client:
yield http_client
async def _render_chat(client) -> dict:
"""Render a minimal chat request and return the GenerateRequest dict."""
resp = await client.post(
"/v1/chat/completions/render",
json={
"model": MODEL_NAME,
"messages": [{"role": "user", "content": "Hello"}],
},
)
assert resp.status_code == 200
return resp.json()
@pytest.mark.asyncio
async def test_streaming_completion_derender_roundtrip(client):
"""Streaming completions derender: chunked text == non streaming text."""
gen_req = await _render_chat(client)
token_ids: list[int] = gen_req["token_ids"][:8]
mid = len(token_ids) // 2
chunk1_ids, chunk2_ids = token_ids[:mid], token_ids[mid:]
# Non streaming baseline.
non_stream_resp = await client.post(
"/v1/completions/derender",
json={
"model": MODEL_NAME,
"generate_responses": [
{
"request_id": "test-ns",
"choices": [
{
"index": 0,
"token_ids": token_ids,
"finish_reason": "stop",
}
],
}
],
},
)
assert non_stream_resp.status_code == 200
expected_text = non_stream_resp.json()["choices"][0]["text"]
# Streaming call 1.
r1 = await client.post(
"/v1/completions/derender",
json={
"stream": True,
"model": MODEL_NAME,
"generate_chunk": {
"request_id": "test-s",
"choices": [
{"index": 0, "token_ids": chunk1_ids, "finish_reason": None}
],
},
"stream_state": None,
},
)
assert r1.status_code == 200
d1 = r1.json()
text1 = d1["chunk"]["choices"][0]["text"]
state1 = d1["stream_state"]
# Streaming call 2 (final chunk).
r2 = await client.post(
"/v1/completions/derender",
json={
"stream": True,
"model": MODEL_NAME,
"generate_chunk": {
"request_id": "test-s",
"choices": [
{"index": 0, "token_ids": chunk2_ids, "finish_reason": "stop"}
],
},
"stream_state": state1,
},
)
assert r2.status_code == 200
text2 = r2.json()["chunk"]["choices"][0]["text"]
assert text1 + text2 == expected_text
@pytest.mark.asyncio
async def test_streaming_chat_derender_roundtrip(client):
"""Streaming chat derender (plain detok): chunked text == non streaming text."""
gen_req = await _render_chat(client)
token_ids: list[int] = gen_req["token_ids"][:8]
mid = len(token_ids) // 2
chunk1_ids, chunk2_ids = token_ids[:mid], token_ids[mid:]
# Non streaming baseline.
ns = await client.post(
"/v1/chat/completions/derender",
json={
"model": MODEL_NAME,
"generate_response": {
"request_id": "test-ns",
"choices": [
{
"index": 0,
"token_ids": token_ids,
"finish_reason": "stop",
}
],
},
},
)
assert ns.status_code == 200
expected_content = ns.json()["choices"][0]["message"]["content"]
# Streaming call 1.
r1 = await client.post(
"/v1/chat/completions/derender",
json={
"stream": True,
"model": MODEL_NAME,
"generate_chunk": {
"request_id": "test-s",
"choices": [
{"index": 0, "token_ids": chunk1_ids, "finish_reason": None}
],
},
"stream_state": None,
},
)
assert r1.status_code == 200
d1 = r1.json()
text1 = d1["chunk"]["choices"][0]["delta"].get("content") or ""
state1 = d1["stream_state"]
# role=assistant on the first chunk
assert d1["chunk"]["choices"][0]["delta"].get("role") == "assistant"
# Streaming call 2.
r2 = await client.post(
"/v1/chat/completions/derender",
json={
"stream": True,
"model": MODEL_NAME,
"generate_chunk": {
"request_id": "test-s",
"choices": [
{"index": 0, "token_ids": chunk2_ids, "finish_reason": "stop"}
],
},
"stream_state": state1,
},
)
assert r2.status_code == 200
d2 = r2.json()
text2 = d2["chunk"]["choices"][0]["delta"].get("content") or ""
# role must NOT be repeated on subsequent chunks
assert d2["chunk"]["choices"][0]["delta"].get("role") is None
assert text1 + text2 == expected_content
@pytest.mark.asyncio
async def test_streaming_derender_invalid_body_returns_400(client):
"""Missing required field in streaming request returns 400."""
r = await client.post(
"/v1/completions/derender",
json={
"stream": True,
# missing required 'model' and 'generate_chunk'
},
)
assert r.status_code == 400
@pytest.mark.asyncio
async def test_streaming_derender_non_object_body_returns_400(client):
"""A non object JSON body (e.g. a list) returns 400, not a 500."""
r = await client.post(
"/v1/completions/derender",
json=[1, 2, 3],
)
assert r.status_code == 400
@pytest.mark.asyncio
async def test_streaming_usage_chunk(client):
"""Usage only final chunk is forwarded with correct token counts."""
gen_req = await _render_chat(client)
token_ids: list[int] = gen_req["token_ids"][:6]
state: dict = {}
# Send content chunk first.
r1 = await client.post(
"/v1/completions/derender",
json={
"stream": True,
"model": MODEL_NAME,
"generate_chunk": {
"request_id": "usage-test",
"choices": [
{"index": 0, "token_ids": token_ids, "finish_reason": "stop"}
],
},
"stream_state": None,
},
)
assert r1.status_code == 200
state = r1.json()["stream_state"]
# Send usage only final chunk.
r2 = await client.post(
"/v1/completions/derender",
json={
"stream": True,
"model": MODEL_NAME,
"generate_chunk": {
"request_id": "usage-test",
"choices": [],
"usage": {
"prompt_tokens": 10,
"completion_tokens": len(token_ids),
"total_tokens": 10 + len(token_ids),
},
},
"stream_state": state,
"prompt_tokens": 10,
},
)
assert r2.status_code == 200
d2 = r2.json()
assert d2["chunk"]["choices"] == []
assert d2["chunk"]["usage"]["prompt_tokens"] == 10
assert d2["chunk"]["usage"]["completion_tokens"] == len(token_ids)
@@ -5,15 +5,17 @@ from http import HTTPStatus
from fastapi import APIRouter, Depends, Request
from fastapi.responses import JSONResponse
from vllm.entrypoints.openai.chat_completion.protocol import ChatCompletionResponse
from vllm.entrypoints.openai.completion.protocol import CompletionResponse
from vllm.entrypoints.openai.engine.protocol import ErrorResponse
from vllm.entrypoints.serve.utils.api_utils import validate_json_request
from vllm.logger import init_logger
from ..token_in_token_out.protocol import (
DerenderChatRequest,
DerenderCompletionRequest,
DerenderChatRequestUnion,
DerenderChatStreamRequest,
DerenderChatStreamResponse,
DerenderCompletionRequestUnion,
DerenderCompletionStreamRequest,
DerenderCompletionStreamResponse,
)
from .serving import ServingDerender
@@ -29,46 +31,96 @@ def derender(request: Request) -> ServingDerender | None:
@router.post(
"/v1/chat/completions/derender",
dependencies=[Depends(validate_json_request)],
response_model=ChatCompletionResponse,
responses={
HTTPStatus.BAD_REQUEST.value: {"model": ErrorResponse},
HTTPStatus.NOT_FOUND.value: {"model": ErrorResponse},
HTTPStatus.INTERNAL_SERVER_ERROR.value: {"model": ErrorResponse},
},
)
async def derender_chat_completion(request: DerenderChatRequest, raw_request: Request):
async def derender_chat_completion(
request: DerenderChatRequestUnion,
raw_request: Request,
):
"""Derender a generate response into a ChatCompletionResponse.
Accepts both non-streaming (``stream=false``, default) and streaming
(``stream=true``) request bodies on the same path; FastAPI validates and
routes on the ``stream`` discriminator.
Non-streaming: body is ``DerenderChatRequest`` (``generate_response`` with
the complete token list). Returns a ``ChatCompletionResponse``.
Streaming: body is ``DerenderChatStreamRequest`` (one ``generate_chunk``
delta + optional ``stream_state``). Returns a ``DerenderChatStreamResponse``
(``chunk`` + ``stream_state``). The client carries ``stream_state`` between
successive calls, one per SSE chunk from ``/inference/v1/generate``.
"""
handler = derender(raw_request)
if handler is None:
raise NotImplementedError(
"The model does not support Chat Completions Derender API"
)
result = await handler.derender_chat_response(request)
if isinstance(request, DerenderChatStreamRequest):
stream_result = await handler.derender_chat_stream_response(request)
if isinstance(stream_result, ErrorResponse):
return JSONResponse(
content=stream_result.model_dump(),
status_code=stream_result.error.code,
)
chunk, stream_state = stream_result
response = DerenderChatStreamResponse(chunk=chunk, stream_state=stream_state)
return JSONResponse(content=response.model_dump())
result = await handler.derender_chat_response(request)
if isinstance(result, ErrorResponse):
return JSONResponse(content=result.model_dump(), status_code=result.error.code)
return JSONResponse(content=result.model_dump())
@router.post(
"/v1/completions/derender",
dependencies=[Depends(validate_json_request)],
response_model=CompletionResponse,
responses={
HTTPStatus.BAD_REQUEST.value: {"model": ErrorResponse},
HTTPStatus.NOT_FOUND.value: {"model": ErrorResponse},
HTTPStatus.INTERNAL_SERVER_ERROR.value: {"model": ErrorResponse},
},
)
async def derender_completion(request: DerenderCompletionRequest, raw_request: Request):
async def derender_completion(
request: DerenderCompletionRequestUnion,
raw_request: Request,
):
"""Derender a generate response into a CompletionResponse.
Accepts both non-streaming (``stream=false``, default) and streaming
(``stream=true``) request bodies on the same path.
Non-streaming: body is ``DerenderCompletionRequest``. Returns a
``CompletionResponse``.
Streaming: body is ``DerenderCompletionStreamRequest`` (one
``generate_chunk`` + optional ``stream_state``). Returns a
``DerenderCompletionStreamResponse`` (``chunk`` + ``stream_state``).
"""
handler = derender(raw_request)
if handler is None:
raise NotImplementedError("The model does not support Completions Derender API")
result = await handler.derender_completion_response(request)
if isinstance(request, DerenderCompletionStreamRequest):
stream_result = await handler.derender_completion_stream_response(request)
if isinstance(stream_result, ErrorResponse):
return JSONResponse(
content=stream_result.model_dump(),
status_code=stream_result.error.code,
)
chunk, stream_state = stream_result
response = DerenderCompletionStreamResponse(
chunk=chunk, stream_state=stream_state
)
return JSONResponse(content=response.model_dump())
result = await handler.derender_completion_response(request)
if isinstance(result, ErrorResponse):
return JSONResponse(content=result.model_dump(), status_code=result.error.code)
return JSONResponse(content=result.model_dump())
+98 -3
View File
@@ -4,8 +4,14 @@ 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.chat_completion.protocol import (
ChatCompletionResponse,
ChatCompletionStreamResponse,
)
from vllm.entrypoints.openai.completion.protocol import (
CompletionResponse,
CompletionStreamResponse,
)
from vllm.entrypoints.openai.engine.protocol import (
ErrorResponse,
UsageInfo,
@@ -28,7 +34,10 @@ from vllm.renderers.online_derenderer import OnlineDerenderer
from ..token_in_token_out.mm_serde import encode_mm_kwargs_item
from ..token_in_token_out.protocol import (
DerenderChatRequest,
DerenderChatStreamRequest,
DerenderCompletionRequest,
DerenderCompletionStreamRequest,
DerenderStreamState,
GenerateResponse,
MultiModalFeatures,
PlaceholderRangeInfo,
@@ -200,7 +209,9 @@ class ServingDerender(BaseServing):
total_prompt_tokens,
total_completion_tokens,
) = await self.online_derenderer.derender_completion(
request.generate_responses, request.prompt_tokens
request.generate_responses,
request.prompt_tokens,
completion_request=request.completion_request,
)
first = request.generate_responses[0]
@@ -237,6 +248,90 @@ class ServingDerender(BaseServing):
kv_transfer_params=kv_params,
)
async def derender_chat_stream_response(
self,
request: DerenderChatStreamRequest,
) -> tuple[ChatCompletionStreamResponse, DerenderStreamState] | ErrorResponse:
"""Streaming counterpart to ``derender_chat_response``.
Processes one ``GenerateStreamResponse`` chunk and returns the
derendered chunk together with the updated client carried state.
``parser is None`` or no ``chat_request`` until reasoning/tool call
functionality added in future PR.
"""
error_check_ret = await self._check_model(request)
if error_check_ret is not None:
return error_check_ret
try:
chunk, updated_state = await self.online_derenderer.derender_chat_stream(
model=request.model,
generate_chunk=request.generate_chunk,
state=request.stream_state,
chat_request=request.chat_request,
prompt_tokens=request.prompt_tokens,
)
except NotImplementedError as exc:
return self.create_error_response(exc)
except ValueError as exc:
return self.create_error_response(str(exc))
except (KeyError, IndexError) as exc:
return self.create_error_response(
f"invalid stream_state: detokenization failed ({exc!r})"
)
logger.debug(
"derender_chat_stream request_id=%s model=%s delta_tokens=%d",
request.generate_chunk.request_id,
request.model,
sum(
len(c.token_ids) for c in request.generate_chunk.choices if c.token_ids
),
)
return chunk, updated_state
async def derender_completion_stream_response(
self,
request: DerenderCompletionStreamRequest,
) -> tuple[CompletionStreamResponse, DerenderStreamState] | ErrorResponse:
"""Streaming counterpart to ``derender_completion_response``.
Processes one ``GenerateStreamResponse`` chunk (one output sequence's
delta) and returns the derendered chunk and updated state.
"""
error_check_ret = await self._check_model(request)
if error_check_ret is not None:
return error_check_ret
try:
(
chunk,
updated_state,
) = await self.online_derenderer.derender_completion_stream(
model=request.model,
generate_chunk=request.generate_chunk,
state=request.stream_state,
prompt_tokens=request.prompt_tokens,
completion_request=request.completion_request,
)
except ValueError as exc:
return self.create_error_response(str(exc))
except (KeyError, IndexError) as exc:
return self.create_error_response(
f"invalid stream_state: detokenization failed ({exc!r})"
)
logger.debug(
"derender_completion_stream request_id=%s model=%s delta_tokens=%d",
request.generate_chunk.request_id,
request.model,
sum(
len(c.token_ids) for c in request.generate_chunk.choices if c.token_ids
),
)
return chunk, updated_state
@staticmethod
def _extract_mm_features(
engine_input: EngineInput,
@@ -1,6 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from typing import Any
from typing import Any, Literal, TypeAlias
from pydantic import (
BaseModel,
@@ -14,8 +14,12 @@ from vllm.config import ModelConfig
from vllm.entrypoints.openai.chat_completion.protocol import (
ChatCompletionLogProbs,
ChatCompletionRequest,
ChatCompletionStreamResponse,
)
from vllm.entrypoints.openai.completion.protocol import (
CompletionRequest,
CompletionStreamResponse,
)
from vllm.entrypoints.openai.completion.protocol import CompletionRequest
from vllm.entrypoints.openai.engine.protocol import StreamOptions, UsageInfo
from vllm.logprobs import Logprob
from vllm.renderers import TokenizeParams
@@ -252,21 +256,16 @@ class GenerateResponse(BaseModel):
)
####### Derender (postprocessing) #######
class DerenderChatRequest(BaseModel):
"""Request for the /v1/chat/completions/derender endpoint (non-streaming).
Wraps a complete GenerateResponse and caller-supplied metadata needed to
produce a fully-formed ChatCompletionResponse without a GPU.
Streaming derender would require a separate endpoint design with
incremental token delivery, ``OutputProcessor``-based detokenization,
and ``parser.parse_delta()`` instead of ``parser.parse()``.
Wraps a complete GenerateResponse and caller supplied metadata needed to
produce a fully formed ChatCompletionResponse without a GPU.
"""
# --8<-- [start:derender-chat-request]
stream: Literal[False] = False
model: str
"""Served model name."""
@@ -299,6 +298,8 @@ class DerenderCompletionRequest(BaseModel):
"""
# --8<-- [start:derender-completion-request]
stream: Literal[False] = False
model: str
"""Served model name."""
@@ -330,3 +331,156 @@ class DerenderCompletionRequest(BaseModel):
f"generate_responses length ({len(self.generate_responses)})"
)
return self
class DerenderStreamState(BaseModel):
"""Per sequence state for stateless streaming derender.
The client carries this between successive per chunk HTTP calls to the
streaming derender endpoint. All fields are plain JSON serializable data.
No opaque tokenizer or parser internals are stored here.
The detokenization strategy carries the incremental decode offsets
directly rather than re-sending the whole token history each chunk.
``detokenize_incrementally`` only ever reads the trailing token window
``prev_tokens[prefix_offset:]``, so we carry just that tail plus the two
offsets. Each chunk resumes exactly where the last one stopped, including
any partially processed multi-byte character (tracked by ``read_offset``),
then trims and rebases the window so it never grows with generation length.
Performance:
- Compute per chunk is O(delta). One ``detokenize_incrementally`` call per
new token, independent of how many tokens preceded it.
- Transport per chunk is O(window). The carried tail is bounded by the
incremental detokenization offset, so cumulative bytes over the wire are
O(n) rather than the O(n^2) a full history round trip would incur.
"""
prev_tokens: list[str] = Field(default_factory=list)
"""Trailing decode window. Token strings from ``prefix_offset`` onward.
Bounded, trimmed and rebased each chunk to the tail
``detokenize_incrementally`` still reads, so it does not grow with the
number of chunks.
"""
prefix_offset: int = Field(default=0, ge=0)
"""Prefix offset into ``prev_tokens`` for incremental detokenization."""
read_offset: int = Field(default=0, ge=0)
"""Read offset into ``prev_tokens`` for incremental detokenization."""
@field_validator("prev_tokens")
@classmethod
def _bound_prev_tokens(cls, v: list[str]) -> list[str]:
# INITIAL_INCREMENTAL_DETOKENIZATION_OFFSET is small (5) and the trimmed
# window is O(offset). A generous limit rejects unusually large or malformed
# payloads without restricting legitimate multi-byte sequences.
limit = 1024
if len(v) > limit:
raise ValueError(f"prev_tokens length ({len(v)}) exceeds maximum ({limit})")
return v
role_sent: bool = False
"""True once the initial ``role: "assistant"`` delta has been emitted.
Prevents re-emitting the role on subsequent chunks even when the detok
window is transiently empty (e.g. usage only final chunk).
"""
# TODO: Properties used in follow on PR for tool call parsing
last_content: str | None = None
"""Last emitted cumulative assistant content text."""
last_reasoning: str | None = None
"""Last emitted cumulative reasoning text."""
last_tool_call_ids: list[str] = Field(default_factory=list)
"""Stable tool-call IDs, assigned once when each call first appears.
Prevents ID regeneration across re-parsing.
"""
class DerenderChatStreamRequest(BaseModel):
"""One chunk streaming derender request for /v1/chat/completions/derender.
The client sends one request per SSE chunk received from
``/inference/v1/generate``. Each request carries the generate chunk
plus the ``stream_state`` returned by the previous call (``None`` on the
first call). The response contains the derendered chunk and the updated
state to be passed to the next call.
This implements stateless no server side session. All mutable state lives in
the client carried ``stream_state``.
"""
stream: Literal[True]
model: str
generate_chunk: GenerateStreamResponse
"""One SSE chunk from ``/inference/v1/generate`` (``stream=True``)."""
stream_state: DerenderStreamState | None = None
"""Client carried detok state from the previous call. ``None`` on first."""
prompt_tokens: int | None = None
"""Prompt token count for usage. Forwarded from the render step."""
chat_request: ChatCompletionRequest | None = None
"""The original (post adjust_request) ChatCompletionRequest from /render."""
class DerenderCompletionStreamRequest(BaseModel):
"""One chunk streaming derender request for /v1/completions/derender.
Parallel to ``DerenderChatStreamRequest`` for the completions endpoint.
Each call processes one SSE chunk (one output sequence's delta) and
returns the derendered chunk plus updated state.
"""
stream: Literal[True]
model: str
generate_chunk: GenerateStreamResponse
"""One SSE chunk from ``/inference/v1/generate``."""
stream_state: DerenderStreamState | None = None
"""Client-carried detok state. ``None`` on the first call."""
prompt_tokens: int | None = None
"""Prompt token count for usage."""
completion_request: CompletionRequest | None = None
"""The original (post adjust_request) CompletionRequest from /render."""
class DerenderChatStreamResponse(BaseModel):
"""Response for one streaming chat derender chunk.
Pairs the derendered SSE chunk with the updated client carried state to
pass to the next call.
"""
chunk: ChatCompletionStreamResponse
stream_state: DerenderStreamState
class DerenderCompletionStreamResponse(BaseModel):
"""Response for one streaming completions derender chunk.
Parallel to ``DerenderChatStreamResponse`` for the completions endpoint.
"""
chunk: CompletionStreamResponse
stream_state: DerenderStreamState
# Determines the type by checking the ``stream`` field's literal value. A body without
# ``stream`` validates as the non-streaming member
# (``stream`` defaults to ``False`` there), so FastAPI can validate and dispatch both
# shapes on a single path.
DerenderChatRequestUnion: TypeAlias = DerenderChatRequest | DerenderChatStreamRequest
DerenderCompletionRequestUnion: TypeAlias = (
DerenderCompletionRequest | DerenderCompletionStreamRequest
)
+4
View File
@@ -17,7 +17,9 @@ from vllm.entrypoints.openai.completion.protocol import (
from vllm.entrypoints.openai.responses.protocol import ResponsesRequest
from vllm.entrypoints.scale_out.token_in_token_out.protocol import (
DerenderChatRequest,
DerenderChatStreamRequest,
DerenderCompletionRequest,
DerenderCompletionStreamRequest,
GenerateRequest,
GenerateResponse,
)
@@ -54,6 +56,7 @@ CompletionLikeRequest: TypeAlias = (
| TokenizeCompletionRequest
| DetokenizeRequest
| DerenderCompletionRequest
| DerenderCompletionStreamRequest
)
ChatLikeRequest: TypeAlias = (
@@ -61,6 +64,7 @@ ChatLikeRequest: TypeAlias = (
| BatchChatCompletionRequest
| TokenizeChatRequest
| DerenderChatRequest
| DerenderChatStreamRequest
)
SpeechToTextRequest: TypeAlias = TranscriptionRequest | TranslationRequest
+281 -6
View File
@@ -10,19 +10,29 @@ from vllm.entrypoints.openai.chat_completion.protocol import (
ChatCompletionNamedToolChoiceParam,
ChatCompletionRequest,
ChatCompletionResponseChoice,
ChatCompletionResponseStreamChoice,
ChatCompletionStreamResponse,
ChatMessage,
)
from vllm.entrypoints.openai.completion.protocol import (
CompletionLogProbs,
CompletionRequest,
CompletionResponseChoice,
CompletionResponseStreamChoice,
CompletionStreamResponse,
)
from vllm.entrypoints.openai.engine.protocol import DeltaMessage, ToolCall, UsageInfo
from vllm.entrypoints.scale_out.token_in_token_out.protocol import (
DerenderStreamState,
GenerateResponse,
GenerateStreamResponse,
)
from vllm.entrypoints.openai.engine.protocol import ToolCall
from vllm.entrypoints.scale_out.token_in_token_out.protocol import GenerateResponse
from vllm.entrypoints.serve.utils.request_logger import RequestLogger
from vllm.logger import init_logger
from vllm.parser import Parser, ParserManager
from vllm.renderers import BaseRenderer
from vllm.tokenizers import TokenizerLike
from vllm.tokenizers.detokenizer_utils import detokenize_incrementally
from vllm.utils import random_uuid
from vllm.utils.async_utils import make_async
@@ -168,9 +178,15 @@ class OnlineDerenderer:
tool_calls=tc_items,
)
else:
# No parser: plain detokenization.
# No parser: plain detokenization honouring the request's
# skip_special_tokens (default True when no request was given).
skip_special = (
chat_request.skip_special_tokens
if chat_request is not None
else True
)
decoded_text = tokenizer.decode(
choice.token_ids, skip_special_tokens=True
choice.token_ids, skip_special_tokens=skip_special
)
message = ChatMessage(role="assistant", content=decoded_text)
@@ -185,23 +201,200 @@ class OnlineDerenderer:
return choices
def _detokenize_delta(
self,
tokenizer: TokenizerLike,
delta_token_ids: list[int],
state: DerenderStreamState,
skip_special_tokens: bool = True,
spaces_between_special_tokens: bool = True,
) -> tuple[str, DerenderStreamState]:
"""Incrementally detokenize ``delta_token_ids`` from prior stream state.
Resumes decoding from the offsets carried in ``state`` rather than
replaying token history. ``state.prev_tokens`` holds the trailing decode
window (from ``prefix_offset`` onward) that ``detokenize_incrementally``
still needs to reproduce any partially read multi-byte character
(tracked by ``read_offset``). The delta tokens are fed straight onto it.
The window is bounded. ``detokenize_incrementally`` never reads before
``prefix_offset``, so after processing we trim ``prev_tokens`` to that
tail and rebase the offsets to it. State transport therefore stays
O(window) per chunk instead of re-sending the full token history.
Args:
tokenizer: The tokenizer to decode with.
delta_token_ids: New token IDs from this generate chunk.
state: Client carried detok state from the previous call.
skip_special_tokens: Passed through to the tokenizer.
spaces_between_special_tokens: Passed through to the tokenizer.
Returns:
(new_text, updated_state) — the delta text for this chunk and the
state to pass to the next call.
"""
prev_tokens = list(state.prev_tokens)
prefix_offset = state.prefix_offset
read_offset = state.read_offset
text_parts: list[str] = []
for tok_id in delta_token_ids:
# prev_tokens is a (possibly empty) list, never None, so this
# always takes the non first iter path and only consumes
# all_input_ids[-1].
new_toks, text, prefix_offset, read_offset = detokenize_incrementally(
tokenizer=tokenizer,
all_input_ids=[tok_id],
prev_tokens=prev_tokens,
prefix_offset=prefix_offset,
read_offset=read_offset,
skip_special_tokens=skip_special_tokens,
spaces_between_special_tokens=spaces_between_special_tokens,
)
prev_tokens = prev_tokens + new_toks
text_parts.append(text)
# Trim to the tail still readable by detokenize_incrementally
# (everything before prefix_offset is dead) and rebase the offsets so
# the carried window stays bounded regardless of generation length.
trimmed = prev_tokens[prefix_offset:]
updated_state = state.model_copy(
update={
"prev_tokens": trimmed,
"prefix_offset": 0,
"read_offset": read_offset - prefix_offset,
}
)
return "".join(text_parts), updated_state
async def derender_chat_stream(
self,
model: str,
generate_chunk: GenerateStreamResponse,
state: DerenderStreamState | None = None,
chat_request: ChatCompletionRequest | None = None,
prompt_tokens: int | None = None,
) -> tuple[ChatCompletionStreamResponse, DerenderStreamState]:
"""Process one GenerateStreamResponse chunk for streaming chat derender.
TODO: parse path for reasoning and tool calls is implemented in future PR.
Unlike OpenAI's API, which always emits ``role: "assistant"`` on the
very first chunk, this emits it on the first chunk with a non empty
``choices`` list. A leading usage only chunk therefore defers the
role to the following content chunk instead of sending an empty
role only delta.
Args:
model: Model name for the response object.
generate_chunk: One SSE chunk from ``/inference/v1/generate``.
state: Client carried detok state (``None`` for first call).
chat_request: Original ChatCompletionRequest from ``/render``.
prompt_tokens: Prompt token count for the usage chunk.
Returns:
(chunk, updated_state) — the derendered SSE chunk and the state
the client must pass to the next call.
"""
if state is None:
state = DerenderStreamState()
if self.parser is not None:
# TODO: Follow on PR will implement the parse path. Check on the
# parser alone (fail closed). A parser configured model must never
# fall through to plain detok on the streaming path, even when
# ``chat_request`` is omitted or reasoning/tool markup would leak
# into ``delta.content``.
raise NotImplementedError(
"Streaming chat derender is not yet supported for models with "
"a reasoning or tool parser configured. Use the non-streaming "
"derender endpoint (stream=false) for parsed output."
)
# A single DerenderStreamState is threaded through every choice in
# this chunk. Correct only when there is at most one choice per SSE
# event (n=1, one call per index), as the streaming derender
# protocol assumes. Multiple choices sharing one chunk would corrupt
# each other's detok window.
if len(generate_chunk.choices) > 1:
raise ValueError(
"derender_chat_stream expects at most one choice per chunk"
)
tokenizer = self.renderer.get_tokenizer()
skip_special = (
chat_request.skip_special_tokens if chat_request is not None else True
)
stream_choices: list[ChatCompletionResponseStreamChoice] = []
updated_state = state
for choice in generate_chunk.choices:
delta_tids = choice.token_ids or []
new_text, updated_state = self._detokenize_delta(
tokenizer, delta_tids, updated_state, skip_special_tokens=skip_special
)
include_role = not updated_state.role_sent
if include_role:
updated_state = updated_state.model_copy(update={"role_sent": True})
delta = DeltaMessage(
role="assistant" if include_role else None,
content=new_text if new_text else None,
)
stream_choices.append(
ChatCompletionResponseStreamChoice(
index=choice.index,
delta=delta,
finish_reason=choice.finish_reason,
)
)
usage: UsageInfo | None = None
if generate_chunk.usage is not None:
u = generate_chunk.usage
pt = prompt_tokens if prompt_tokens is not None else (u.prompt_tokens or 0)
ct = u.completion_tokens or 0
usage = UsageInfo(
prompt_tokens=pt,
completion_tokens=ct,
total_tokens=pt + ct,
)
chunk = ChatCompletionStreamResponse(
id=generate_chunk.request_id,
model=model,
choices=stream_choices,
usage=usage,
)
return chunk, updated_state
async def derender_completion(
self,
generate_responses: list[GenerateResponse],
prompt_tokens: list[int] | None = None,
completion_request: CompletionRequest | None = None,
) -> tuple[list[CompletionResponseChoice], int, int]:
return await self._derender_completion_async(generate_responses, prompt_tokens)
return await self._derender_completion_async(
generate_responses, prompt_tokens, completion_request
)
def _derender_completion(
self,
generate_responses: list[GenerateResponse],
prompt_tokens: list[int] | None = None,
completion_request: CompletionRequest | None = None,
) -> tuple[list[CompletionResponseChoice], int, int]:
n = len(generate_responses)
prompt_tokens_list: list[int] = (
prompt_tokens if prompt_tokens is not None else [0] * n
)
skip_special = (
completion_request.skip_special_tokens
if completion_request is not None
else True
)
tokenizer = self.renderer.get_tokenizer()
choices: list[CompletionResponseChoice] = []
total_prompt_tokens = 0
@@ -217,7 +410,7 @@ class OnlineDerenderer:
)
decoded_text = tokenizer.decode(
choice.token_ids, skip_special_tokens=True
choice.token_ids, skip_special_tokens=skip_special
)
completion_logprobs = None
if choice.logprobs is not None:
@@ -239,6 +432,88 @@ class OnlineDerenderer:
return choices, total_prompt_tokens, total_completion_tokens
async def derender_completion_stream(
self,
model: str,
generate_chunk: GenerateStreamResponse,
state: DerenderStreamState | None = None,
prompt_tokens: int | None = None,
completion_request: CompletionRequest | None = None,
) -> tuple[CompletionStreamResponse, DerenderStreamState]:
"""Process one GenerateStreamResponse chunk for streaming completions.
Each call takes one SSE chunk from ``/inference/v1/generate`` plus the
client carried ``stream_state`` and returns a ``CompletionStreamResponse``
chunk and the updated state.
The generate stream emits one choice per SSE event, so this method
processes one output sequence at a time. For ``n > 1`` the client
maintains one ``DerenderStreamState`` per ``choice.index``.
Args:
model: Model name for the response object.
generate_chunk: One SSE chunk from ``/inference/v1/generate``.
state: Client carried detok state (``None`` → first call).
prompt_tokens: Prompt token count for usage (from the render step).
completion_request: Original CompletionRequest from ``/render``;
supplies ``skip_special_tokens``.
Returns:
(chunk, updated_state) — the derendered chunk and updated state.
"""
if state is None:
state = DerenderStreamState()
# See the equivalent check in derender_chat_stream: a single
# DerenderStreamState is threaded through every choice in this
# chunk, so more than one choice per chunk would corrupt the
# detok window across choices.
if len(generate_chunk.choices) > 1:
raise ValueError(
"derender_completion_stream expects at most one choice per chunk"
)
tokenizer = self.renderer.get_tokenizer()
skip_special = (
completion_request.skip_special_tokens
if completion_request is not None
else True
)
stream_choices: list[CompletionResponseStreamChoice] = []
updated_state = state
for choice in generate_chunk.choices:
delta_tids = choice.token_ids or []
new_text, updated_state = self._detokenize_delta(
tokenizer, delta_tids, updated_state, skip_special_tokens=skip_special
)
stream_choices.append(
CompletionResponseStreamChoice(
index=choice.index,
text=new_text,
finish_reason=choice.finish_reason,
)
)
usage: UsageInfo | None = None
if generate_chunk.usage is not None:
u = generate_chunk.usage
pt = prompt_tokens if prompt_tokens is not None else (u.prompt_tokens or 0)
ct = u.completion_tokens or 0
usage = UsageInfo(
prompt_tokens=pt,
completion_tokens=ct,
total_tokens=pt + ct,
)
chunk = CompletionStreamResponse(
id=generate_chunk.request_id,
model=model,
choices=stream_choices,
usage=usage,
)
return chunk, updated_state
def _parse_token_id_placeholder(token: str) -> int | None:
"""Extract token ID from a 'token_id:N' placeholder string."""