diff --git a/tests/entrypoints/scale_out/derender/test_derender_stream.py b/tests/entrypoints/scale_out/derender/test_derender_stream.py new file mode 100644 index 00000000000..014a8d7bc6b --- /dev/null +++ b/tests/entrypoints/scale_out/derender/test_derender_stream.py @@ -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) diff --git a/vllm/entrypoints/scale_out/derender/api_router.py b/vllm/entrypoints/scale_out/derender/api_router.py index 3f88d51f0a9..4a139144db9 100644 --- a/vllm/entrypoints/scale_out/derender/api_router.py +++ b/vllm/entrypoints/scale_out/derender/api_router.py @@ -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()) diff --git a/vllm/entrypoints/scale_out/derender/serving.py b/vllm/entrypoints/scale_out/derender/serving.py index 613ff65ad0e..4b72ce2e9fc 100644 --- a/vllm/entrypoints/scale_out/derender/serving.py +++ b/vllm/entrypoints/scale_out/derender/serving.py @@ -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, 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 c22e70b014c..bf387809670 100644 --- a/vllm/entrypoints/scale_out/token_in_token_out/protocol.py +++ b/vllm/entrypoints/scale_out/token_in_token_out/protocol.py @@ -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 +) diff --git a/vllm/entrypoints/serve/engine/typing.py b/vllm/entrypoints/serve/engine/typing.py index 2e01c092c7b..a6878d50015 100644 --- a/vllm/entrypoints/serve/engine/typing.py +++ b/vllm/entrypoints/serve/engine/typing.py @@ -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 diff --git a/vllm/renderers/online_derenderer.py b/vllm/renderers/online_derenderer.py index 22e0b45790e..a1285d8e178 100644 --- a/vllm/renderers/online_derenderer.py +++ b/vllm/renderers/online_derenderer.py @@ -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."""