From cabb4b77a2c9590aa2ce01c3dce556c89ba87b6c Mon Sep 17 00:00:00 2001 From: Mergen Nachin Date: Tue, 15 Sep 2026 12:27:32 -0700 Subject: [PATCH 1/2] [llm_server] Fix transcript replay and reasoning response handling Fall back to rendered text when the assistant boundary cannot be verified. Compare supplied reasoning against the finalized client-visible response before replaying stored tokens, preserving omitted reasoning while invalidating explicit edits. Discard whitespace-only Muse Glimmer blocks and reject non-boolean return_reasoning values before generation. Document the response contract and verify complete and streaming turns with full prompt assertions. Run serving tests for relevant pull requests, including the Muse adapter tests. Make tokenizer paths explicit and cover exact BPE prompt assembly with a portable fixture. --- .github/workflows/_llm_server.yml | 7 +- .github/workflows/pull.yml | 22 ++ .github/workflows/trunk.yml | 20 +- examples/llm_server/python/README.md | 5 +- .../llm_server/python/openai_transcript.py | 53 +-- examples/llm_server/python/protocol.py | 1 + examples/llm_server/python/serving_chat.py | 19 +- examples/llm_server/python/tests/conftest.py | 2 + .../llm_server/python/tests/test_contract.py | 76 ++++ .../python/tests/test_tool_calls.py | 47 --- .../python/tests/test_warm_resume_scaffold.py | 328 +++++++++++++++--- examples/llm_server/spec/README.md | 17 + examples/models/muse-glimmer/README.md | 9 +- examples/models/muse-glimmer/serving/serve.py | 13 +- .../models/muse-glimmer/tests/test_serve.py | 125 +++++-- 15 files changed, 575 insertions(+), 169 deletions(-) diff --git a/.github/workflows/_llm_server.yml b/.github/workflows/_llm_server.yml index 6492a3106ee..12d04423112 100644 --- a/.github/workflows/_llm_server.yml +++ b/.github/workflows/_llm_server.yml @@ -36,12 +36,15 @@ jobs: python -m pip install --progress-bar off \ -r examples/llm_server/python/requirements.txt \ httpx \ - pytest + pytest \ + tokenizers export PYTHONPATH="$(dirname "${PWD}"):${PYTHONPATH:-}" export PYTHONDONTWRITEBYTECODE=1 - python -m pytest -q examples/llm_server/python/tests + python -m pytest -q \ + examples/llm_server/python/tests \ + examples/models/muse-glimmer/tests/test_serve.py cmake -S . -B cmake-out \ -DCMAKE_BUILD_TYPE=Release \ diff --git a/.github/workflows/pull.yml b/.github/workflows/pull.yml index e51901411ab..26326f1c21a 100644 --- a/.github/workflows/pull.yml +++ b/.github/workflows/pull.yml @@ -752,6 +752,28 @@ jobs: id-token: write contents: read + llm-server: + needs: [changed-files, run-decision] + if: | + github.event_name == 'pull_request' && ( + contains(needs.changed-files.outputs.changed-files, 'examples/llm_server') || + contains(needs.changed-files.outputs.changed-files, 'examples/models/muse-glimmer/serving') || + contains(needs.changed-files.outputs.changed-files, 'examples/models/muse-glimmer/tests') || + contains(needs.changed-files.outputs.changed-files, 'examples/models/muse_glimmer') || + contains(needs.changed-files.outputs.changed-files, 'extension/llm/runner') || + contains(needs.changed-files.outputs.changed-files, 'extension/llm/tokenizers') || + contains(needs.changed-files.outputs.changed-files, '.github/workflows/_get-changed-files.yml') || + contains(needs.changed-files.outputs.changed-files, '.github/workflows/_ci-run-decision.yml') || + contains(needs.changed-files.outputs.changed-files, '.github/workflows/_llm_server.yml') || + contains(needs.changed-files.outputs.changed-files, '.github/workflows/pull.yml') || + contains(needs.changed-files.outputs.changed-files, '.github/workflows/trunk.yml') || + needs.run-decision.outputs.is-full-run == 'true' + ) + uses: ./.github/workflows/_llm_server.yml + permissions: + id-token: write + contents: read + unittest: uses: ./.github/workflows/_unittest.yml permissions: diff --git a/.github/workflows/trunk.yml b/.github/workflows/trunk.yml index 1b3749fdebf..234c035880b 100644 --- a/.github/workflows/trunk.yml +++ b/.github/workflows/trunk.yml @@ -994,12 +994,20 @@ jobs: llm-server: needs: [changed-files, run-decision] if: | - contains(needs.changed-files.outputs.changed-files, 'examples/llm_server') || - contains(needs.changed-files.outputs.changed-files, 'extension/llm/runner') || - contains(needs.changed-files.outputs.changed-files, 'extension/llm/tokenizers') || - contains(needs.changed-files.outputs.changed-files, '.github/workflows/_llm_server.yml') || - contains(needs.changed-files.outputs.changed-files, '.github/workflows/trunk.yml') || - needs.run-decision.outputs.is-full-run == 'true' + github.event_name != 'pull_request' && ( + contains(needs.changed-files.outputs.changed-files, 'examples/llm_server') || + contains(needs.changed-files.outputs.changed-files, 'examples/models/muse-glimmer/serving') || + contains(needs.changed-files.outputs.changed-files, 'examples/models/muse-glimmer/tests') || + contains(needs.changed-files.outputs.changed-files, 'examples/models/muse_glimmer') || + contains(needs.changed-files.outputs.changed-files, 'extension/llm/runner') || + contains(needs.changed-files.outputs.changed-files, 'extension/llm/tokenizers') || + contains(needs.changed-files.outputs.changed-files, '.github/workflows/_get-changed-files.yml') || + contains(needs.changed-files.outputs.changed-files, '.github/workflows/_ci-run-decision.yml') || + contains(needs.changed-files.outputs.changed-files, '.github/workflows/_llm_server.yml') || + contains(needs.changed-files.outputs.changed-files, '.github/workflows/pull.yml') || + contains(needs.changed-files.outputs.changed-files, '.github/workflows/trunk.yml') || + needs.run-decision.outputs.is-full-run == 'true' + ) uses: ./.github/workflows/_llm_server.yml permissions: id-token: write diff --git a/examples/llm_server/python/README.md b/examples/llm_server/python/README.md index f9020a7b342..f6f9d2271fe 100644 --- a/examples/llm_server/python/README.md +++ b/examples/llm_server/python/README.md @@ -114,7 +114,7 @@ Two layers, both contract-focused (assert on the wire, not internals): ```bash # 1. Model-free tests — unit coverage plus loopback disconnect integration. -pip install pytest httpx +pip install pytest httpx tokenizers pytest tests/ # 2. Conformance — black-box, against a LIVE server (real model, or llama.cpp/mlx-lm). @@ -126,6 +126,9 @@ real server/protocol/streaming code is tested over HTTP without a `.pte`. The worker JSONL protocol is covered separately by `tests/test_worker_client.py`, and `tests/test_stream_disconnect.py` uses real loopback Uvicorn/TCP plus a model-free subprocess to verify disconnect cancellation end to end. +The BPE splice tests use an in-memory tokenizer with no model downloads. Optional +integration tests use local tokenizer directories set with `QWEN_HF_DIR`, +`GEMMA_HF_DIR`, or `MUSE_GLIMMER_HF_DIR` and require `transformers`. ## Architecture diff --git a/examples/llm_server/python/openai_transcript.py b/examples/llm_server/python/openai_transcript.py index a6bf5b75e6d..b238968faf3 100644 --- a/examples/llm_server/python/openai_transcript.py +++ b/examples/llm_server/python/openai_transcript.py @@ -16,10 +16,10 @@ token ids and a fingerprint of the response. On the next request each prior assistant turn is replaced with a sentinel, the conversation is rendered once, and the rendered text is split on the sentinels with the stored ids spliced back -in -- only for turns whose fingerprint matches the incoming message (an edited, -branched, or reused history is never substituted with stale ids) and whose ids -are present (a stop-trimmed turn is left as text). The worker's exact-token -prefix check is the final backstop. +in -- only for turns whose content/tool fingerprint and any supplied reasoning +match the recorded response, and whose ids are present (a stop-trimmed turn is +left as text). An edit invalidates that turn and all later records. The worker's +exact-token prefix check separately protects KV reuse. """ import hashlib @@ -111,8 +111,7 @@ def __init__(self, template: ChatTemplate): # boundary verification. _header_fn = getattr(template, "assistant_header", None) self._assist_hdr = _header_fn() if _header_fn else _ASSIST_HDR - # session_id -> [{"fp": str, "ids": list[int] | None}, ...] (one per - # assistant turn we produced, in order). Cleared on reset/close. + # One record per generated assistant turn, cleared on reset/close. self._turns: dict[str, list[dict]] = {} @staticmethod @@ -135,19 +134,24 @@ def _assistant_fingerprint(content, tool_calls) -> str: blob = json.dumps([content or "", norm], sort_keys=True, ensure_ascii=False) return hashlib.sha1(blob.encode("utf-8")).hexdigest() + @staticmethod + def _reasoning_fingerprint(reasoning_content: Optional[str]) -> Optional[bytes]: + if reasoning_content is None: + return None + return hashlib.sha256(reasoning_content.encode("utf-8")).digest() + def _normalize_scaffold(self, text_chunk: str, preamble: str) -> Optional[str]: """Force the scaffold region (between the last assistant header in `text_chunk` and its end) to equal `preamble`, so the worker re-tokenizes the exact resident scaffold. The region is empty (history stripped it -> insert) or a think scaffold (history preserved it -> replace). Returns the adjusted text, or None if it isn't a recognized scaffold (-> text fallback).""" + if not self._assist_hdr: + return None h = text_chunk.rfind(self._assist_hdr) if h == -1: - # No assistant header: with a scaffold to reproduce this is - # ambiguous (-> text fallback); without one there is nothing to - # normalize, so splicing still works for templates with a different - # assistant header. - return None if preamble else text_chunk + # Without a verified boundary, splicing can duplicate template framing. + return None base = h + len(self._assist_hdr) if not preamble: # No generation scaffold: the worker prefills nothing ahead of the @@ -266,9 +270,9 @@ def build_prompt_input( """Return a PromptInput: token-ID segments when this session has faithful stored ids for matching prior assistant turns, else the plain rendered text. Each incoming assistant turn is matched IN ORDER against the stored - records and only spliced when (a) its fingerprint matches what we returned - (else the history diverged -> stop, splice nothing further) and (b) we - kept faithful ids for it (a stop-trimmed turn's None -> rendered as text). + records and only spliced when its content/tool calls and any supplied + reasoning match what we returned, and we kept faithful ids for it. Omitted + reasoning permits reuse; an explicit edit invalidates the stored tail. Falls back to text on a sentinel collision or a render that dropped/duplicated a sentinel.""" stored = self._turns.get(session_id or "") @@ -286,13 +290,18 @@ def build_prompt_input( if k >= len(stored): break m = messages[pos] - if self._assistant_fingerprint(m.content, m.tool_calls) != stored[k]["fp"]: + record = stored[k] + if self._assistant_fingerprint(m.content, m.tool_calls) != record["fp"] or ( + "reasoning_content" in m.model_fields_set + and self._reasoning_fingerprint(m.reasoning_content) + != record["reasoning_fp"] + ): diverged_at = k # this stored turn and every later one are stale break - if stored[k]["ids"] is not None: + if record["ids"] is not None: splice[pos] = { - "ids": stored[k]["ids"], - "preamble": stored[k].get("preamble", ""), + "ids": record["ids"], + "preamble": record.get("preamble", ""), } if diverged_at is not None: # Drop the stale tail from the first mismatch so an edited/branched @@ -351,14 +360,17 @@ def record_assistant_turn( generated_token_ids: list, prior_turns: int, preamble: str = "", + reasoning_content: Optional[str] = None, ) -> None: """Record this turn's {fingerprint, generated ids, generation preamble} at `prior_turns` (the assistant-turn count of the request it answers). Records at/after that index are dropped first, so a regenerated/branched turn replaces stale records rather than shadowing later hits. ids is None when the worker omitted them (stop-trimmed -> non-resumable), kept for - positional alignment. `preamble` is the generation scaffold (e.g. the - Qwen3 `` block) reproduced ahead of the spliced ids next request.""" + positional alignment. `reasoning_content` is the client-visible value, + including None when the client opted out. `preamble` is the generation + scaffold (e.g. the Qwen3 `` block) reproduced ahead of the spliced + ids next request.""" if not session_id: return turns = self._turns.setdefault(session_id, []) @@ -366,6 +378,7 @@ def record_assistant_turn( turns.append( { "fp": self._assistant_fingerprint(content, tool_calls), + "reasoning_fp": self._reasoning_fingerprint(reasoning_content), "ids": list(generated_token_ids) if generated_token_ids else None, "preamble": preamble, } diff --git a/examples/llm_server/python/protocol.py b/examples/llm_server/python/protocol.py index 0530719e922..069ace32b31 100644 --- a/examples/llm_server/python/protocol.py +++ b/examples/llm_server/python/protocol.py @@ -42,6 +42,7 @@ class ChatMessage(BaseModel): # ResponseMessage.reasoning_content so a multi-turn client can echo an assistant # turn's reasoning back into the request; a chat template that renders prior # reasoning needs it here, and without the field it is dropped at parse. + # Omission allows stored-token replay; an explicit edit invalidates that turn. reasoning_content: Optional[str] = None tool_calls: Optional[list[ToolCall]] = None tool_call_id: Optional[str] = None diff --git a/examples/llm_server/python/serving_chat.py b/examples/llm_server/python/serving_chat.py index 61cb671fee3..36c21795aec 100644 --- a/examples/llm_server/python/serving_chat.py +++ b/examples/llm_server/python/serving_chat.py @@ -143,8 +143,7 @@ def _return_reasoning(req: ChatCompletionRequest) -> bool: # return them unless the client explicitly opts out, matching # SGLang/llama.cpp. Explicit {"return_reasoning": False} opts out. kwargs = req.chat_template_kwargs or {} - value = kwargs.get("return_reasoning", True) - return value if isinstance(value, bool) else False + return kwargs.get("return_reasoning", True) @staticmethod def _to_openai_tool_call(item: ToolCallItem) -> ToolCall: @@ -352,8 +351,16 @@ def _finish_reason( @staticmethod def _reject_invalid_values(req: ChatCompletionRequest) -> None: - """Reject out-of-range values (invalid_value); these take precedence over + """Reject invalid types/ranges (invalid_value); these take precedence over the unsupported-parameter error.""" + template_kwargs = req.chat_template_kwargs or {} + if not isinstance(template_kwargs.get("return_reasoning", True), bool): + raise APIError( + 400, + "chat_template_kwargs.return_reasoning must be a boolean.", + "invalid_request_error", + "invalid_value", + ) if req.temperature is not None and ( not math.isfinite(req.temperature) or req.temperature < 0.0 @@ -553,13 +560,12 @@ async def _complete( tool_calls, reasoning, content = self._extract_response( req, self._truncate_raw(text, req) ) - # Record after the response is finalized: the fingerprint is of exactly - # what we return (content + tool_calls), so the next turn can confirm the - # client echoed this turn before splicing its ids. + # Compare future echoes against the finalized client-visible response. self._transcript.record_assistant_turn( session_id=req.session_id, content=content, tool_calls=tool_calls, + reasoning_content=reasoning, generated_token_ids=stats.generated_token_ids, prior_turns=sum(1 for m in req.messages if m.role == "assistant"), preamble=preamble, @@ -745,6 +751,7 @@ def chunk(delta: DeltaMessage, finish=None) -> str: session_id=req.session_id, content=content, tool_calls=tool_calls, + reasoning_content=reasoning or None, generated_token_ids=stats.generated_token_ids, prior_turns=sum(1 for m in req.messages if m.role == "assistant"), preamble=preamble, diff --git a/examples/llm_server/python/tests/conftest.py b/examples/llm_server/python/tests/conftest.py index 097b12eb99a..40b68f11ac4 100644 --- a/examples/llm_server/python/tests/conftest.py +++ b/examples/llm_server/python/tests/conftest.py @@ -128,6 +128,7 @@ def _make( max_named_sessions=0, gen_ids=None, reuse=0, + reasoning_extractor=None, ): fake = FakeRunner( tokens, @@ -147,6 +148,7 @@ def _make( "test-model", max_context=max_context, tool_detector_cls=HermesDetector, + reasoning_extractor=reasoning_extractor, ) return TestClient(build_app(serving, "test-model")), fake diff --git a/examples/llm_server/python/tests/test_contract.py b/examples/llm_server/python/tests/test_contract.py index eb45d8e4681..a90415e276a 100644 --- a/examples/llm_server/python/tests/test_contract.py +++ b/examples/llm_server/python/tests/test_contract.py @@ -134,6 +134,82 @@ def test_chat_streaming_protocol(make_client): assert chunks[-1]["choices"][0]["finish_reason"] == "stop" +@pytest.mark.parametrize("stream", [False, True]) +@pytest.mark.parametrize( + "options,returns_reasoning", + [ + ({}, True), + ({"chat_template_kwargs": {}}, True), + ({"chat_template_kwargs": {"return_reasoning": True}}, True), + ({"chat_template_kwargs": {"return_reasoning": False}}, False), + ], + ids=["default", "empty-kwargs", "enabled", "disabled"], +) +def test_reasoning_response_contract(make_client, stream, options, returns_reasoning): + def split_reasoning(text): + reasoning, _, content = text.partition("ANSWER:") + return reasoning, content + + client, _ = make_client( + tokens=["my plan", "ANSWER:", "visible text"], + reasoning_extractor=split_reasoning, + ) + response = client.post( + "/v1/chat/completions", + json={ + "model": "test-model", + "messages": [{"role": "user", "content": "hi"}], + "stream": stream, + **options, + }, + ) + assert response.status_code == 200 + if stream: + assert response.headers["content-type"].startswith("text/event-stream") + chunks, done = _sse_chunks(response.text) + assert done + messages = [chunk["choices"][0]["delta"] for chunk in chunks] + assert chunks[-1]["choices"][0]["finish_reason"] == "stop" + else: + choice = response.json()["choices"][0] + messages = [choice["message"]] + assert choice["finish_reason"] == "stop" + assert "".join(message.get("content", "") for message in messages) == "visible text" + if returns_reasoning: + assert ( + "".join(message.get("reasoning_content", "") for message in messages) + == "my plan" + ) + else: + assert all("reasoning_content" not in message for message in messages) + + +@pytest.mark.parametrize("stream", [False, True]) +@pytest.mark.parametrize("value", ["true", "false", 0, 1, 1.0, None, [], {}]) +def test_return_reasoning_rejects_non_boolean_before_generation( + make_client, stream, value +): + client, fake = make_client(max_named_sessions=1) + response = client.post( + "/v1/chat/completions", + json={ + "model": "test-model", + "messages": [{"role": "user", "content": "hi"}], + "session_id": "validation", + "stream": stream, + "chat_template_kwargs": {"return_reasoning": value}, + }, + ) + assert response.status_code == 400 + assert response.headers["content-type"].startswith("application/json") + error = response.json()["error"] + assert error["type"] == "invalid_request_error" + assert error["code"] == "invalid_value" + assert "chat_template_kwargs.return_reasoning" in error["message"] + assert fake.opened_log == [] + assert fake.captured_config is None + + def test_request_params_forwarded_to_generation(make_client): # Contract behavior: the server must honor all supported sampling controls. client, fake = make_client() diff --git a/examples/llm_server/python/tests/test_tool_calls.py b/examples/llm_server/python/tests/test_tool_calls.py index d2867738ed9..86ca778bb81 100644 --- a/examples/llm_server/python/tests/test_tool_calls.py +++ b/examples/llm_server/python/tests/test_tool_calls.py @@ -219,50 +219,3 @@ def test_parallel_calls_in_one_message(make_client): ).json() calls = body["choices"][0]["message"]["tool_calls"] assert [json.loads(c["function"]["arguments"])["city"] for c in calls] == ["A", "B"] - - -# --- reasoning_content default ------------------------------------------- - - -def _reasoning_serving(): - from executorch.examples.llm_server.python.serving_chat import ServingChat - - def _split(text): - head, _, tail = text.partition("ANSWER:") - return head.strip() or None, tail.strip() - - class _FakeTemplate: - def turn_stop_sequences(self): - return [] - - def special_tokens(self): - return [] - - return ServingChat(None, _FakeTemplate(), "m", reasoning_extractor=_split) - - -def _reasoning_req(**kw): - from executorch.examples.llm_server.python.protocol import ChatCompletionRequest - - return ChatCompletionRequest(messages=[{"role": "user", "content": "hi"}], **kw) - - -def test_reasoning_returned_by_default(): - serving = _reasoning_serving() - calls, reasoning, content = serving._extract_response( - _reasoning_req(), "my plan ANSWER: visible text" - ) - assert calls is None - assert reasoning == "my plan" - assert content == "visible text" - - -def test_reasoning_explicit_opt_out(): - serving = _reasoning_serving() - calls, reasoning, content = serving._extract_response( - _reasoning_req(chat_template_kwargs={"return_reasoning": False}), - "my plan ANSWER: visible text", - ) - assert calls is None - assert reasoning is None - assert content == "visible text" diff --git a/examples/llm_server/python/tests/test_warm_resume_scaffold.py b/examples/llm_server/python/tests/test_warm_resume_scaffold.py index 86c4e42b451..4b91c45e4f1 100644 --- a/examples/llm_server/python/tests/test_warm_resume_scaffold.py +++ b/examples/llm_server/python/tests/test_warm_resume_scaffold.py @@ -28,6 +28,7 @@ FunctionCall, ToolCall, ) +from tokenizers import decoders, models, Tokenizer HDR = "<|im_start|>assistant\n" NOTHINK = "\n\n\n\n" # no-think generation preamble / preserved block @@ -88,6 +89,9 @@ class _FakeOtherHeader: OHDR = "<|start_header_id|>assistant<|end_header_id|>\n\n" + def assistant_header(self): + return self.OHDR + def render(self, messages, tools=None, template_kwargs=None): out = [] for m in messages: @@ -103,8 +107,11 @@ def render(self, messages, tools=None, template_kwargs=None): class _FakeGemma: + def __init__(self, header=GEMMA_HDR): + self._header = header + def assistant_header(self): - return GEMMA_HDR + return self._header def render(self, messages, tools=None, template_kwargs=None): out = [""] @@ -237,15 +244,17 @@ def test_no_scaffold_template_is_unchanged(): def test_non_qwen_header_no_scaffold_still_splices(): - # Regression: a no-scaffold template whose assistant header isn't the - # Qwen/ChatML one must still get token-id splicing (the normalization is a - # no-op when preamble == "", not a hard requirement for the Qwen header). - st = OpenAITranscriptState(_FakeOtherHeader()) + # A correctly configured non-ChatML header still supports splicing. + fake = _FakeOtherHeader() + st = OpenAITranscriptState(fake) + enc = _ByteTokenizer.encode + gen_ids = enc("a1") + resident = enc(fake.render(_msgs(("user", "u1")))) + gen_ids st.record_assistant_turn( session_id="s", content="a1", tool_calls=None, - generated_token_ids=[9, 9], + generated_token_ids=gen_ids, prior_turns=0, preamble="", ) @@ -257,8 +266,11 @@ def test_non_qwen_header_no_scaffold_still_splices(): tools=None, template_kwargs=None, ) - assert pi.segments is not None # splicing NOT disabled by the missing header - assert any(s.get("ids") == [9, 9] for s in pi.segments) # ids actually spliced + assert pi.segments is not None + assembled = _assemble(pi.segments, enc) + assert assembled == enc(fake.render(msgs)) + assert assembled[: len(resident)] == resident + assert any(s.get("ids") == gen_ids for s in pi.segments) def test_custom_assistant_header_inserts_scaffold(): @@ -453,8 +465,12 @@ def test_tool_turn_splices_despite_reserialized_args(): assert any(s.get("ids") == [1, 2, 3] for s in pi.segments) -def test_gemma_tool_span_ignores_close_marker_inside_string(): - st = OpenAITranscriptState(_FakeGemma()) +@pytest.mark.parametrize("header", [GEMMA_HDR, HDR], ids=["matched", "mismatched"]) +def test_gemma_tool_span_ignores_close_marker_inside_string(header): + st = OpenAITranscriptState(_FakeGemma(header)) + enc = _ByteTokenizer.encode + raw = '<|tool_call>call:bash{command:<|"|>printf ok<|"|>}' + gen_ids = enc(raw) call = ToolCall( id="call_1", function=FunctionCall( @@ -465,7 +481,7 @@ def test_gemma_tool_span_ignores_close_marker_inside_string(): session_id="s", content="", tool_calls=[call], - generated_token_ids=[101, 102, 103], + generated_token_ids=gen_ids, prior_turns=0, preamble="", ) @@ -474,12 +490,9 @@ def test_gemma_tool_span_ignores_close_marker_inside_string(): ChatMessage(role="assistant", content="", tool_calls=[call]), ChatMessage(role="tool", tool_call_id="call_1", content="done"), ] - rendered = ( - "<|turn>user\nu1\n" - "<|turn>model\n" - '<|tool_call>call:bash{command:<|"|>printf ok<|"|>}' - "<|tool_response>done" - ) + first_prompt = "<|turn>user\nu1\n<|turn>model\n" + trailing = "<|tool_response>done" + rendered = first_prompt + raw + trailing pi = st.build_prompt_input( session_id="s", messages=msgs, @@ -487,23 +500,26 @@ def test_gemma_tool_span_ignores_close_marker_inside_string(): tools=None, template_kwargs=None, ) + assembled = enc(pi.text) if pi.segments is None else _assemble_ids(pi.segments) + assert assembled == enc(first_prompt) + gen_ids + enc(trailing) + if header != GEMMA_HDR: + assert pi.segments is None + assert pi.text == rendered + return assert pi.segments is not None - assert any(s.get("ids") == [101, 102, 103] for s in pi.segments) + assert any(s.get("ids") == gen_ids for s in pi.segments) suffix = "".join( - s.get("text", "") - for s in pi.segments[_ids_index(pi.segments, [101, 102, 103]) + 1 :] + s.get("text", "") for s in pi.segments[_ids_index(pi.segments, gen_ids) + 1 :] ) assert suffix == "<|tool_response>done" # --- 5b. Token-level fidelity against the real tokenizer (gated/skipped) ----- -_MODEL = os.environ.get( - "QWEN_HF_DIR", "/home/mnachin/local/scripts/models/Qwen3.5-35B-A3B-HQQ-INT4" -) -_HAVE_MODEL = os.path.isdir(_MODEL) +_MODEL = os.environ.get("QWEN_HF_DIR", "") +_HAVE_MODEL = bool(_MODEL) and os.path.isdir(_MODEL) _skip = pytest.mark.skipif( - not _HAVE_MODEL, reason=f"real Qwen tokenizer dir not present: {_MODEL}" + not _HAVE_MODEL, reason="set QWEN_HF_DIR to a local Qwen tokenizer directory" ) @@ -526,12 +542,10 @@ def _assemble(segs, enc): return out -_GEMMA_MODEL = os.environ.get( - "GEMMA_HF_DIR", "/home/mnachin/local/scripts/models/gemma-4-31B-it-HQQ-INT4" -) -_HAVE_GEMMA = os.path.isdir(_GEMMA_MODEL) +_GEMMA_MODEL = os.environ.get("GEMMA_HF_DIR", "") +_HAVE_GEMMA = bool(_GEMMA_MODEL) and os.path.isdir(_GEMMA_MODEL) _skip_gemma = pytest.mark.skipif( - not _HAVE_GEMMA, reason=f"real Gemma tokenizer dir not present: {_GEMMA_MODEL}" + not _HAVE_GEMMA, reason="set GEMMA_HF_DIR to a local Gemma tokenizer directory" ) @@ -792,6 +806,7 @@ def test_thinking_turn_splice_reproduces_resident_prefix(): generated_token_ids=gen_ids, prior_turns=0, preamble="", + reasoning_content="\nthink\n", ) msgs = [ ChatMessage(role="user", content="u1"), @@ -849,6 +864,188 @@ def test_plain_turn_splice_reproduces_resident_prefix(): assert trailing.startswith("<|eot|>") +@pytest.mark.parametrize( + "recorded_reasoning, echoed_fields, matches", + [ + pytest.param("ORIGINAL", {}, True, id="omitted"), + pytest.param( + "ORIGINAL", {"reasoning_content": "ORIGINAL"}, True, id="unchanged" + ), + pytest.param("ORIGINAL", {"reasoning_content": "EDITED"}, False, id="changed"), + pytest.param( + "ORIGINAL", {"reasoning_content": "ORIGINAL "}, False, id="whitespace" + ), + pytest.param("ORIGINAL", {"reasoning_content": ""}, False, id="empty"), + pytest.param("ORIGINAL", {"reasoning_content": None}, False, id="null"), + pytest.param(None, {"reasoning_content": None}, True, id="unchanged-null"), + pytest.param(None, {"reasoning_content": ""}, False, id="null-to-empty"), + ], +) +def test_reasoning_echo_must_match_when_explicit( + recorded_reasoning, echoed_fields, matches +): + fake = _FakeHarmony() + st = OpenAITranscriptState(fake) + raw = " to=user<|message|>a1" + if recorded_reasoning is not None: + raw = ( + " to=self<|message|>" + recorded_reasoning + "<|eom|>" + "<|start|>assistant" + raw + ) + resident, gen_ids = _harmony_resident(fake, raw) + st.record_assistant_turn( + session_id="s", + content="a1", + tool_calls=None, + generated_token_ids=gen_ids, + prior_turns=0, + preamble="", + reasoning_content=recorded_reasoning, + ) + msgs = [ + ChatMessage(role="user", content="u1"), + ChatMessage(role="assistant", content="a1", **echoed_fields), + ChatMessage(role="user", content="u2"), + ] + rendered = fake.render(msgs) + pi = st.build_prompt_input( + session_id="s", + messages=msgs, + rendered_prompt=rendered, + tools=None, + template_kwargs=None, + ) + assembled = ( + _ByteTokenizer.encode(pi.text) + if pi.segments is None + else _assemble_ids(pi.segments) + ) + if matches: + assert pi.segments is not None + trailing = "<|eot|><|start|>user<|message|>u2<|eot|><|start|>assistant" + assert assembled == resident + _ByteTokenizer.encode(trailing) + else: + assert assembled == _ByteTokenizer.encode(rendered) + assert pi.segments is None + assert pi.text == rendered + + +def test_reasoning_edit_invalidates_stored_tail_without_later_resurrection(): + fake = _FakeHarmony() + st = OpenAITranscriptState(fake) + enc = _ByteTokenizer.encode + msgs = [] + generated = [] + for i in range(1, 4): + reasoning = f"reasoning {i}" + raw = ( + " to=self<|message|>" + reasoning + "<|eom|>" + f"<|start|>assistant to=user<|message|>a{i}" + ) + gen_ids = enc(raw) + generated.append(gen_ids) + st.record_assistant_turn( + session_id="s", + content=f"a{i}", + tool_calls=None, + generated_token_ids=gen_ids, + prior_turns=i - 1, + preamble="", + reasoning_content=reasoning, + ) + msgs.extend( + [ + ChatMessage(role="user", content=f"u{i}"), + ChatMessage( + role="assistant", content=f"a{i}", reasoning_content=reasoning + ), + ] + ) + msgs.append(ChatMessage(role="user", content="u4")) + original = msgs[3] + msgs[3] = original.model_copy(update={"reasoning_content": "EDITED"}) + for second_turn in (msgs[3], original): + msgs[3] = second_turn + rendered = fake.render(msgs) + pi = st.build_prompt_input( + session_id="s", + messages=msgs, + rendered_prompt=rendered, + tools=None, + template_kwargs=None, + ) + assert pi.segments is not None + assert _assemble_ids(pi.segments) == enc(rendered) + assert [s["ids"] for s in pi.segments if "ids" in s] == [generated[0]] + + +@pytest.mark.parametrize("thinking", [False, True]) +def test_harmony_bpe_splice_preserves_resident_ids_and_full_prompt(thinking): + # A tiny real BPE with a merge crossing the prompt/generation boundary. + # Assets and training are unnecessary; every token and merge is explicit. + specials = ["", "<|start|>", "<|message|>", "<|eom|>", "<|eot|>"] + vocab = {chr(i): i for i in range(128)} + vocab.update({token: len(vocab) + i for i, token in enumerate(specials)}) + vocab["t "] = len(vocab) + vocab["ab"] = len(vocab) + tokenizer = Tokenizer(models.BPE(vocab, merges=[("t", " "), ("a", "b")])) + tokenizer.add_special_tokens(specials) + tokenizer.decoder = decoders.Fuse() + + def enc(text): + return tokenizer.encode(text, add_special_tokens=False).ids + + fake = _FakeHarmony() + first_prompt = fake.render(_msgs(("user", "u1"))) + reasoning = "\nthink\n" if thinking else None + raw_prefix = " to=user<|message|>" + if thinking: + raw_prefix = ( + " to=self<|message|>" + reasoning + "<|eom|>" + "<|start|>assistant" + raw_prefix + ) + raw = raw_prefix + "ab" + # Valid generated tokens need not use the encoder's canonical segmentation. + # These decode faithfully, but re-encoding would merge the final a + b. + gen_ids = enc(raw_prefix) + [vocab["a"], vocab["b"]] + assert tokenizer.decode(gen_ids, skip_special_tokens=False) == raw + assert gen_ids != enc(raw) + assert vocab["<|eot|>"] not in gen_ids + assert enc(first_prompt + " ") != enc(first_prompt) + enc(" ") + resident = enc(first_prompt) + gen_ids + + st = OpenAITranscriptState(fake) + st.record_assistant_turn( + session_id="s", + content="ab", + tool_calls=None, + generated_token_ids=gen_ids, + prior_turns=0, + preamble="", + reasoning_content=reasoning, + ) + msgs = [ + ChatMessage(role="user", content="u1"), + ChatMessage(role="assistant", content="ab", reasoning_content=reasoning), + ChatMessage(role="user", content="u2"), + ] + pi = st.build_prompt_input( + session_id="s", + messages=msgs, + rendered_prompt=fake.render(msgs), + tools=None, + template_kwargs=None, + ) + assert pi.segments is not None + trailing = "<|eot|><|start|>user<|message|>u2<|eot|><|start|>assistant" + assembled = _assemble(pi.segments, enc) + assert assembled[: len(resident)] == resident + assert assembled == resident + enc(trailing) + full_text = first_prompt + raw + trailing + assert tokenizer.decode(assembled, skip_special_tokens=False) == full_text + assert enc(full_text) != assembled # Text equality cannot establish reuse. + + class _FakeMismatchedHeader: """Mirrors the production failure mode: the adapter provides assistant_header() (production adapters always do) but it returns the @@ -860,8 +1057,11 @@ class _FakeMismatchedHeader: # The default header: correct for ChatML, wrong for this template. HDR = "<|im_start|>assistant\n" + def __init__(self, header=HDR): + self._header = header + def assistant_header(self): - return self.HDR + return self._header def render(self, messages, tools=None, template_kwargs=None): out = [""] @@ -931,45 +1131,69 @@ def test_user_header_literal_does_not_truncate_conversation(): assert "KEEP" in pi.text -def test_mismatched_header_without_literal_still_splices(): - # Positive control: without a confusing literal, the default header never - # matches (h == -1, nothing to normalize) and splicing proceeds normally - # for the mismatched-header template. - st = OpenAITranscriptState(_FakeMismatchedHeader()) +@pytest.mark.parametrize( + "header", + [HDR, "", _FakeHarmony.HDR], + ids=["mismatched", "empty", "matched"], +) +def test_splice_requires_verified_header_to_preserve_full_prompt(header): + fake = _FakeMismatchedHeader(header) + st = OpenAITranscriptState(fake) + enc = _ByteTokenizer.encode + # Generated IDs include recipient framing, but exclude the terminal EOT. + # Keeping the rendered recipient as well would duplicate model input even + # if the worker's resident-prefix check subsequently chooses a cold prefill. + gen_ids = enc(" to=user<|message|>a1") + resident = enc(fake.render(_msgs(("user", "plain question")))) + gen_ids st.record_assistant_turn( session_id="s", content="a1", tool_calls=None, - generated_token_ids=[5, 6], + generated_token_ids=gen_ids, prior_turns=0, preamble="", ) msgs = _msgs(("user", "plain question"), ("assistant", "a1"), ("user", "u2")) + rendered = fake.render(msgs) pi = st.build_prompt_input( session_id="s", messages=msgs, - rendered_prompt=st._template.render(msgs), + rendered_prompt=rendered, tools=None, template_kwargs=None, ) - assert pi.segments is not None - assert any(s.get("ids") == [5, 6] for s in pi.segments) - - -def test_mistral_bracket_terminator_tail_falls_back(): + assembled = enc(pi.text) if pi.segments is None else _assemble_ids(pi.segments) + trailing = "<|eot|><|start|>user<|message|>u2<|eot|><|start|>assistant" + assert assembled == resident + enc(trailing) + assert assembled == enc(rendered) + if header == _FakeHarmony.HDR: + assert pi.segments is not None + assert any(s.get("ids") == gen_ids for s in pi.segments) + else: + assert pi.segments is None + assert pi.text == rendered + + +@pytest.mark.parametrize( + "u1", ["plain question", "Explain <|im_start|>assistant\nKEEP"] +) +def test_mistral_bracket_terminator_tail_falls_back(u1): # Nemo-style `[/INST]` terminators are turn structure no blocklist can # enumerate exhaustively: the `KEEP[/INST]` tail is short and single-line # yet must fall back, preserving KEEP and the instruction boundary. - st = OpenAITranscriptState(_FakeMistralNemo()) + fake = _FakeMistralNemo() + st = OpenAITranscriptState(fake) + enc = _ByteTokenizer.encode + gen_ids = enc("a1") + resident = enc(fake.render(_msgs(("user", u1)))) + gen_ids st.record_assistant_turn( session_id="s", content="a1", tool_calls=None, - generated_token_ids=[5, 6], + generated_token_ids=gen_ids, prior_turns=0, preamble="", ) - u1 = "Explain <|im_start|>assistant\nKEEP" msgs = _msgs(("user", u1), ("assistant", "a1"), ("user", "u2")) rendered = st._template.render(msgs) pi = st.build_prompt_input( @@ -980,7 +1204,7 @@ def test_mistral_bracket_terminator_tail_falls_back(): template_kwargs=None, ) assert pi.segments is None and pi.text == rendered - assert "KEEP[/INST]" in pi.text + assert enc(pi.text) == resident + enc("[INST]u2[/INST]") # --- generation_preamble threads tools ------------------------------------ @@ -1016,14 +1240,11 @@ def apply_chat_template( # --- Muse Glimmer real-template thinking-turn prefix (gated/skipped) -------- -_GLIMMER_MODEL = os.environ.get( - "GLIMMER_HF_DIR", - "/data/users/mnachin/scripts/agent-runtime-bench/data/executorch/hf", -) -_HAVE_GLIMMER = os.path.isdir(_GLIMMER_MODEL) +_GLIMMER_MODEL = os.environ.get("MUSE_GLIMMER_HF_DIR", "") +_HAVE_GLIMMER = bool(_GLIMMER_MODEL) and os.path.isdir(_GLIMMER_MODEL) _skip_glimmer = pytest.mark.skipif( not _HAVE_GLIMMER, - reason=f"real Muse Glimmer tokenizer dir not present: {_GLIMMER_MODEL}", + reason="set MUSE_GLIMMER_HF_DIR to a local Muse Glimmer tokenizer directory", ) @@ -1071,6 +1292,7 @@ def test_glimmer_real_template_thinking_turn_splice_reproduces_resident(): generated_token_ids=gen_ids, prior_turns=0, preamble=tmpl.generation_preamble(), + reasoning_content=reasoning, ) msgs = [ ChatMessage(role="user", content=u1), diff --git a/examples/llm_server/spec/README.md b/examples/llm_server/spec/README.md index a675a925ee3..0f8bda86e64 100644 --- a/examples/llm_server/spec/README.md +++ b/examples/llm_server/spec/README.md @@ -26,6 +26,11 @@ uses the worker's unset/random value. `model` must match the id returned by `/v1/models`; unknown ids return `404 model_not_found`. +For models with a reasoning extractor, `chat_template_kwargs.return_reasoning` +defaults to `true`. Set it to the boolean `false` to omit reasoning from the +response; this does not disable the model's reasoning computation. Non-boolean +values return `400 invalid_request_error` (`code: "invalid_value"`). + **Rejected** with `400 invalid_request_error` (`code: "unsupported_parameter"`) rather than silently ignored — a client relying on them would otherwise get wrong behavior: `n` (> 1), `reasoning_effort`, @@ -41,6 +46,8 @@ ignored. Non-streaming response: `chat.completion` with one `choice` (`message.role = "assistant"`, string `content` or `tool_calls`, `finish_reason` ∈ `stop` | `length` | `tool_calls`) and a `usage` block. +When reasoning is returned, `message.reasoning_content` contains the extracted +reasoning text separately from visible `content` and `tool_calls`. `usage.prompt_tokens_details.cached_tokens` reports prompt tokens served from the session's resident state instead of prefetched this request (0 when the turn fully prefilled); the streaming usage chunk carries the same field. @@ -50,6 +57,8 @@ first chunk carries `delta.role = "assistant"`, subsequent chunks carry `delta.content` (or buffered `delta.tool_calls`), a final chunk carries `finish_reason`, optionally a usage-only chunk (with `stream_options.include_usage`), terminated by `data: [DONE]`. +For models with a reasoning extractor, output is buffered and returned reasoning +is emitted as `delta.reasoning_content` before visible content or tool calls. ### Tool calling @@ -84,3 +93,11 @@ implemented worker-side for engines that support it — a named session whose ne request is an exact-token extension of its resident context prefills only the new suffix. All KV/resident state lives inside the worker/session, never the control plane. + +For named sessions, an unchanged assistant reply can reuse its original generated +token IDs. Clients may omit `reasoning_content` when echoing a reply. If the field +is supplied, it must match the value returned to that client; a changed value +(including explicit `null` or `""`) invalidates that turn's stored IDs and later +records. The updated history is then rendered normally. If an assistant boundary +cannot be verified, the server also uses the rendered text. The worker checks +the resulting token sequence before reusing KV state in either case. diff --git a/examples/models/muse-glimmer/README.md b/examples/models/muse-glimmer/README.md index fd68bb0c29a..0f2a3e65269 100644 --- a/examples/models/muse-glimmer/README.md +++ b/examples/models/muse-glimmer/README.md @@ -243,10 +243,17 @@ The model id must match `--model-id`. The `compat` entries: - `supportsDeveloperRole` — the template renders no `developer` turn, so pi's system prompt is otherwise dropped silently. - `supportsReasoningEffort` — the server rejects `reasoning_effort` with a 400. -- `return_reasoning` — returns the `to=self` channel as `reasoning_content`. +- `return_reasoning` — defaults to `true`, returning the `to=self` channel as + `reasoning_content` (or `delta.reasoning_content` when streaming). Set + `"chatTemplateKwargs": { "return_reasoning": false }` to omit it. This only + controls the response; the model still computes reasoning. - `sendSessionAffinityHeaders` — optional, for per-conversation sessions; needs `--max-sessions` above 1. Set `contextWindow` to the export's context length (`128K` is 131072) and pass the same value as `--max-context`. Add `"input": ["text", "image"]` for a vision export. + +For direct HTTP requests, the opt-out is +`"chat_template_kwargs": { "return_reasoning": false }`. The flag must be a JSON +boolean; strings, numbers, and `null` return a 400 error. diff --git a/examples/models/muse-glimmer/serving/serve.py b/examples/models/muse-glimmer/serving/serve.py index 89c30e96c62..7ce3122ecf3 100644 --- a/examples/models/muse-glimmer/serving/serve.py +++ b/examples/models/muse-glimmer/serving/serve.py @@ -82,12 +82,11 @@ def _strip_muse_glimmer_header(text: str) -> str: def _extract_muse_glimmer_reasoning(text: str) -> tuple[str | None, str]: """Split a Harmony turn into private `to=self` bodies and visible messages. - The returned reasoning is client-facing text: multiple thinking blocks are - joined with a newline (matching SGLang's reasoning parser) so no protocol - framing leaks to clients. Replay fidelity does NOT come from this text -- - the transcript fingerprint covers content + tool calls (not reasoning), - and warm resume splices the stored generated token ids over whatever - reasoning the client echoes back. + Nonempty thinking bodies retain their whitespace and are joined with a + newline without adding channel framing. Interior text is not otherwise + sanitized. Warm resume compares any supplied reasoning against the returned + text before splicing stored generated token ids, preserving the original + channel framing for unchanged echoes or omitted reasoning. """ matches = list(_MUSE_GLIMMER_ADDRESSED_HEADER_RE.finditer(text)) if not matches: @@ -103,7 +102,7 @@ def _extract_muse_glimmer_reasoning(text: str) -> tuple[str | None, str]: end = matches[index + 1].start() if index + 1 < len(matches) else len(text) body = text[match.end() : end] body = re.sub(r"(?:<\|eom\|>|<\|eot\|>)\s*$", "", body) - if not body: + if not body.strip(): continue if match.group(1) == "self": reasoning.append(body) diff --git a/examples/models/muse-glimmer/tests/test_serve.py b/examples/models/muse-glimmer/tests/test_serve.py index 53509ab8e8e..ef7e50a663f 100644 --- a/examples/models/muse-glimmer/tests/test_serve.py +++ b/examples/models/muse-glimmer/tests/test_serve.py @@ -8,6 +8,7 @@ import asyncio import base64 +import json import pathlib from types import SimpleNamespace @@ -576,6 +577,8 @@ def _benc(text): def _assemble(prompt): + if prompt.text is not None: + return _benc(prompt.text) out = [] for seg in prompt.segments or []: out += _benc(seg["text"]) if "text" in seg else list(seg["ids"]) @@ -592,6 +595,8 @@ def _glimmer_serving(raw_texts, template=None): template or _StubTemplate(), "test-model", reasoning_extractor=serve._extract_muse_glimmer_reasoning, + content_filter=serve._strip_muse_glimmer_header, + content_filter_specials=serve._MUSE_GLIMMER_HEADER_SPECIALS, ) return serving, runtime @@ -618,6 +623,41 @@ def test_thinking_survives_response_verbatim(): assert msg.content == "The answer is 42." +@pytest.mark.parametrize("recipient", ["self", "user"]) +def test_whitespace_only_channel_does_not_add_response_gaps(recipient): + raw = ( + "to=self<|message|>\nThink first.\n<|eom|>" + "<|start|>assistant to=user<|message|>Part one.<|eom|>" + f"<|start|>assistant to={recipient}<|message|> \t\n<|eom|>" + "<|start|>assistant to=user<|message|>Part two.<|eot|>" + ) + serving, _ = _glimmer_serving([{"text": raw}]) + response = asyncio.run( + serving.create( + ChatCompletionRequest(messages=[ChatMessage(role="user", content="hi")]) + ) + ) + message = response.choices[0].message + assert message.reasoning_content == "\nThink first.\n" + assert message.content == "Part one.\n\nPart two." + + +def test_whitespace_only_reasoning_is_omitted(): + raw = ( + "to=self<|message|> \t\n<|eom|>" + "<|start|>assistant to=user<|message|>Done.<|eot|>" + ) + serving, _ = _glimmer_serving([{"text": raw}]) + response = asyncio.run( + serving.create( + ChatCompletionRequest(messages=[ChatMessage(role="user", content="hi")]) + ) + ) + message = response.choices[0].message + assert message.reasoning_content is None + assert message.content == "Done." + + def test_multiple_thinking_blocks_joined_without_protocol_markers(): # Thinking blocks are client-facing text: they must be joined with a # newline (SGLang parity) so no <|eom|>/<|start|>/to=self framing leaks @@ -643,50 +683,83 @@ def test_multiple_thinking_blocks_joined_without_protocol_markers(): assert "<|" not in reasoning and "to=self" not in reasoning -def test_thinking_turn_resumes_exact_prefix_over_two_turns(): - # Extraction + splice composition at the worker seam: turn 2's prompt must - # assemble to the resident prompt plus the trailing turn-2 text exactly, - # so the worker reuses instead of refilling. The recorded ids are the - # worker's non-terminal generated ids (no terminal <|eot|>); the trailing - # text supplies the terminator. Full-prompt equality, not a prefix check. +async def _create_message(serving, request): + response = await serving.create(request) + if not request.stream: + return response.choices[0].message.model_dump(exclude_none=True) + message = {"role": "assistant"} + done = False + async for event in response: + payload = event.removeprefix("data: ").strip() + if payload == "[DONE]": + done = True + continue + chunk = json.loads(payload) + assert "error" not in chunk + delta = chunk["choices"][0]["delta"] + for field in ("content", "reasoning_content"): + if field in delta: + message[field] = message.get(field, "") + delta[field] + assert done + return message + + +@pytest.mark.parametrize("stream", [False, True]) +@pytest.mark.parametrize("echo_mode", ["edited", "unchanged", "omitted", "opt-out"]) +def test_thinking_turn_replays_ids_only_for_unchanged_history(stream, echo_mode): + # The worker reports non-terminal ids; the next rendered turn supplies EOS. + # Two raw thinking blocks become one client-visible reasoning string. raw1 = ( - " to=self<|message|>\nLet me think.\n<|eom|>" + " to=self<|message|>\nFirst step.\n<|eom|>" + "<|start|>assistant to=self<|message|>Second step.<|eom|>" "<|start|>assistant to=user<|message|>156" ) raw2 = " to=user<|message|>210" + template = _HarmonyTemplate() serving, runtime = _glimmer_serving( [ - {"text": raw1 + "<|eot|>", "gen_ids": _benc(raw1)}, - {"text": raw2 + "<|eot|>", "gen_ids": _benc(raw2)}, + {"text": raw1, "gen_ids": _benc(raw1)}, + {"text": raw2, "gen_ids": _benc(raw2)}, ], - template=_HarmonyTemplate(), + template=template, ) u1 = ChatMessage(role="user", content="What is 12*13?") first = asyncio.run( - serving.create(ChatCompletionRequest(messages=[u1], session_id="s")) - ) - echo = ChatMessage( - role="assistant", - content=first.choices[0].message.content, - reasoning_content=first.choices[0].message.reasoning_content, - ) - second = asyncio.run( - serving.create( + _create_message( + serving, ChatCompletionRequest( - messages=[u1, echo, ChatMessage(role="user", content="And?")], + messages=[u1], session_id="s", - ) + stream=stream, + chat_template_kwargs=( + {"return_reasoning": False} if echo_mode == "opt-out" else None + ), + ), ) ) + assert first["content"] == "156" + if echo_mode == "opt-out": + assert "reasoning_content" not in first + else: + assert first["reasoning_content"] == "\nFirst step.\n\nSecond step." + if echo_mode == "edited": + first["reasoning_content"] = "Use this corrected reasoning instead." + elif echo_mode == "omitted": + first.pop("reasoning_content") + request2 = ChatCompletionRequest( + messages=[u1, ChatMessage(**first), ChatMessage(role="user", content="And?")], + session_id="s", + ) + second = asyncio.run(serving.create(request2)) assert second.choices[0].message.content == "210" prompt2 = runtime.prompts[1] - assert prompt2.segments is not None # prior turn spliced as exact ids resident = _benc(runtime.prompts[0].text + raw1) trailing = "<|eot|><|start|>user<|message|>And?<|eot|><|start|>assistant" - assert _assemble(prompt2) == resident + _benc(trailing) - # Realism pins (worker_loop.h: generated ids exclude the terminal EOS). - assert not raw1.endswith("<|eot|>") - assert trailing.startswith("<|eot|>") + if echo_mode == "edited": + assert _assemble(prompt2) == _benc(template.render(request2.messages)) + assert _assemble(prompt2) != resident + _benc(trailing) + else: + assert _assemble(prompt2) == resident + _benc(trailing) def test_extract_muse_glimmer_reasoning_plain_text_fallback(): From cd7b41ff40aabefd04f36f9f862a276fd926dd2b Mon Sep 17 00:00:00 2001 From: Mergen Nachin Date: Thu, 17 Sep 2026 10:14:07 -0700 Subject: [PATCH 2/2] [llm_server] Fix replay recovery and assistant header configuration Preserve assistant-turn indices after history edits, treat null reasoning as omitted, and expose the assistant header with a startup mismatch warning. Clarify return_reasoning validation and response visibility behavior. Validated with 367 passing serving tests (6 optional tokenizer tests skipped), 88 focused HTTP/prompt scenarios, flake8, and ufmt. Authored with Codex. --- examples/llm_server/python/README.md | 8 +++ examples/llm_server/python/chat_template.py | 13 +++- .../llm_server/python/openai_transcript.py | 55 ++++++-------- examples/llm_server/python/protocol.py | 2 +- examples/llm_server/python/server.py | 8 +++ examples/llm_server/python/serving_chat.py | 4 +- .../llm_server/python/tests/test_contract.py | 8 ++- .../python/tests/test_server_launcher.py | 62 ++++++++++++++++ .../llm_server/python/tests/test_sessions.py | 6 +- .../llm_server/python/tests/test_template.py | 19 +++++ .../python/tests/test_warm_resume_scaffold.py | 71 ++++++++++++++++++- examples/llm_server/spec/README.md | 33 ++++++--- examples/models/muse-glimmer/serving/serve.py | 4 +- .../models/muse-glimmer/tests/test_serve.py | 60 +++++++++++++++- 14 files changed, 293 insertions(+), 60 deletions(-) create mode 100644 examples/llm_server/python/tests/test_server_launcher.py diff --git a/examples/llm_server/python/README.md b/examples/llm_server/python/README.md index f6f9d2271fe..d7ea0f927de 100644 --- a/examples/llm_server/python/README.md +++ b/examples/llm_server/python/README.md @@ -60,12 +60,20 @@ Key flags: | Flag | Effect | |------|--------| | `--hf-tokenizer` | model's HF chat template (required unless fallback) | +| `--assistant-header` | exact assistant generation header, including trailing whitespace (default: ChatML) | | `--allow-chatml-fallback` | opt into approximate ChatML when no HF tokenizer | | `--no-think` | default `enable_thinking=False` (e.g. Qwen3) | | `--max-context N` | reject over-long prompts with 400 instead of failing mid-gen | | `--num-runners N` | Worker processes — **1 only** (one worker hosts many isolated sessions on one weight load; more would duplicate weights) | | `--worker-bin PATH` | path to a model worker binary that speaks the llm_server JSONL protocol | +Set `--assistant-header` to the model template's exact generation boundary when +it differs from ChatML. For Llama 3 templates, add +`--assistant-header $'<|start_header_id|>assistant<|end_header_id|>\n\n'` in Bash; +the `$'...'` quoting supplies literal newlines. The launcher warns once at startup +if the configured header is absent from a rendered probe. Unverified boundaries +use the rendered text, which can reduce KV reuse. + ## Smoke test ```bash diff --git a/examples/llm_server/python/chat_template.py b/examples/llm_server/python/chat_template.py index 92d869e6898..a63fbb916ac 100644 --- a/examples/llm_server/python/chat_template.py +++ b/examples/llm_server/python/chat_template.py @@ -106,6 +106,7 @@ def __init__( # chat_template_kwargs override these. self._defaults = default_template_kwargs or {} self._assistant_header = assistant_header + self._warned_assistant_header = False self._strip_rendered_prefix = strip_rendered_prefix self._append_generation_prompt_after_tool_response = ( append_generation_prompt_after_tool_response @@ -229,8 +230,6 @@ def generation_preamble( the resident one (for Qwen3 the scaffold is tool-independent -> same key). Returns ``""`` for the fallback / no-scaffold templates (fix is a no-op). """ - if self._hf is None: - return "" merged = {**self._defaults, **(template_kwargs or {})} if tools: try: @@ -251,7 +250,15 @@ def generation_preamble( template_kwargs=template_kwargs, ) marker = self._assistant_header - idx = rendered.rfind(marker) + idx = rendered.rfind(marker) if marker else -1 + if idx == -1 and not self._warned_assistant_header: + logger.warning( + "Assistant header %r was not found in the chat template's generation " + "prompt. Stored-token replay may fall back to rendered text; configure " + "assistant_header (--assistant-header for the generic server).", + marker, + ) + self._warned_assistant_header = True preamble = rendered[idx + len(marker) :] if idx != -1 else "" self._preamble_cache[key] = preamble return preamble diff --git a/examples/llm_server/python/openai_transcript.py b/examples/llm_server/python/openai_transcript.py index b238968faf3..660efbc0920 100644 --- a/examples/llm_server/python/openai_transcript.py +++ b/examples/llm_server/python/openai_transcript.py @@ -16,7 +16,7 @@ token ids and a fingerprint of the response. On the next request each prior assistant turn is replaced with a sentinel, the conversation is rendered once, and the rendered text is split on the sentinels with the stored ids spliced back -in -- only for turns whose content/tool fingerprint and any supplied reasoning +in -- only for turns whose content/tool fingerprint and any supplied reasoning string match the recorded response, and whose ids are present (a stop-trimmed turn is left as text). An edit invalidates that turn and all later records. The worker's exact-token prefix check separately protects KV reuse. @@ -111,8 +111,8 @@ def __init__(self, template: ChatTemplate): # boundary verification. _header_fn = getattr(template, "assistant_header", None) self._assist_hdr = _header_fn() if _header_fn else _ASSIST_HDR - # One record per generated assistant turn, cleared on reset/close. - self._turns: dict[str, list[dict]] = {} + # Keyed by assistant-turn index, including gaps after invalidation. + self._turns: dict[str, dict[int, dict]] = {} @staticmethod def _assistant_fingerprint(content, tool_calls) -> str: @@ -271,44 +271,36 @@ def build_prompt_input( stored ids for matching prior assistant turns, else the plain rendered text. Each incoming assistant turn is matched IN ORDER against the stored records and only spliced when its content/tool calls and any supplied - reasoning match what we returned, and we kept faithful ids for it. Omitted - reasoning permits reuse; an explicit edit invalidates the stored tail. + reasoning string match what we returned, and we kept faithful ids for it. + Omitted or null reasoning permits reuse; a string edit invalidates the tail. Falls back to text on a sentinel collision or a render that dropped/duplicated a sentinel.""" stored = self._turns.get(session_id or "") if not stored: return PromptInput(text=rendered_prompt) - # Positional: stored[k] is the k-th assistant turn WE generated, matched - # against the k-th assistant message in the request. A client-injected - # turn (few-shot exemplar, pre-seeded turn, reused session) shifts that - # alignment -> fingerprint mismatch at k -> stop splicing. Always safe - # (text fallback + worker prefix backstop); just a lower hit rate. + # Missing records render as text without shifting later turn indices. positions = [i for i, m in enumerate(messages) if m.role == "assistant"] splice: dict[int, dict] = {} # message index -> {"ids", "preamble"} - diverged_at = None for k, pos in enumerate(positions): - if k >= len(stored): - break + record = stored.get(k) + if record is None: + continue m = messages[pos] - record = stored[k] if self._assistant_fingerprint(m.content, m.tool_calls) != record["fp"] or ( - "reasoning_content" in m.model_fields_set + m.reasoning_content is not None and self._reasoning_fingerprint(m.reasoning_content) != record["reasoning_fp"] ): - diverged_at = k # this stored turn and every later one are stale + # Discard the stale tail without shifting subsequent turn indices. + self._turns[session_id or ""] = { + index: record for index, record in stored.items() if index < k + } break if record["ids"] is not None: splice[pos] = { "ids": record["ids"], "preamble": record.get("preamble", ""), } - if diverged_at is not None: - # Drop the stale tail from the first mismatch so an edited/branched - # earlier turn can't shadow future requests; the matched prefix still - # splices, the rest stays text until reset/close. Safe either way: - # stale ids are never spliced and the worker's prefix check backstops. - del stored[diverged_at:] if not splice: return PromptInput(text=rendered_prompt) tool_splice = { @@ -373,16 +365,15 @@ def record_assistant_turn( ids next request.""" if not session_id: return - turns = self._turns.setdefault(session_id, []) - del turns[prior_turns:] - turns.append( - { - "fp": self._assistant_fingerprint(content, tool_calls), - "reasoning_fp": self._reasoning_fingerprint(reasoning_content), - "ids": list(generated_token_ids) if generated_token_ids else None, - "preamble": preamble, - } - ) + turns = self._turns.setdefault(session_id, {}) + for index in [index for index in turns if index >= prior_turns]: + del turns[index] + turns[prior_turns] = { + "fp": self._assistant_fingerprint(content, tool_calls), + "reasoning_fp": self._reasoning_fingerprint(reasoning_content), + "ids": list(generated_token_ids) if generated_token_ids else None, + "preamble": preamble, + } def reset(self, session_id: str) -> None: self._turns.pop(session_id, None) diff --git a/examples/llm_server/python/protocol.py b/examples/llm_server/python/protocol.py index 069ace32b31..bf2d1da17bb 100644 --- a/examples/llm_server/python/protocol.py +++ b/examples/llm_server/python/protocol.py @@ -42,7 +42,7 @@ class ChatMessage(BaseModel): # ResponseMessage.reasoning_content so a multi-turn client can echo an assistant # turn's reasoning back into the request; a chat template that renders prior # reasoning needs it here, and without the field it is dropped at parse. - # Omission allows stored-token replay; an explicit edit invalidates that turn. + # Omission/null allows stored-token replay; a string edit invalidates that turn. reasoning_content: Optional[str] = None tool_calls: Optional[list[ToolCall]] = None tool_call_id: Optional[str] = None diff --git a/examples/llm_server/python/server.py b/examples/llm_server/python/server.py index 8f24b8187f4..534f05c70cb 100644 --- a/examples/llm_server/python/server.py +++ b/examples/llm_server/python/server.py @@ -171,6 +171,12 @@ def main() -> None: help="Allow approximate generic ChatML templating when --hf-tokenizer is absent. " "Off by default: the fallback can't reproduce model-specific controls.", ) + p.add_argument( + "--assistant-header", + default="<|im_start|>assistant\n", + help="Exact template text ending the assistant generation header, including " + "any trailing whitespace. Defaults to the ChatML header.", + ) p.add_argument( "--model-id", default="executorch", help="Model id reported on /v1/models" ) @@ -218,7 +224,9 @@ def main() -> None: args.hf_tokenizer, default_template_kwargs=default_template_kwargs, allow_fallback=args.allow_chatml_fallback, + assistant_header=args.assistant_header, ) + template.generation_preamble() worker = _spawn(args) # one worker hosting many isolated sessions runtime = SessionRuntime(worker) serving = ServingChat( diff --git a/examples/llm_server/python/serving_chat.py b/examples/llm_server/python/serving_chat.py index 36c21795aec..319220ef797 100644 --- a/examples/llm_server/python/serving_chat.py +++ b/examples/llm_server/python/serving_chat.py @@ -139,9 +139,7 @@ def _split_reasoning(self, text: str) -> tuple[Optional[str], str]: @staticmethod def _return_reasoning(req: ChatCompletionRequest) -> bool: - # Default ON: thinking models bill reasoning tokens either way; - # return them unless the client explicitly opts out, matching - # SGLang/llama.cpp. Explicit {"return_reasoning": False} opts out. + # Response visibility only; the model still computes reasoning on opt-out. kwargs = req.chat_template_kwargs or {} return kwargs.get("return_reasoning", True) diff --git a/examples/llm_server/python/tests/test_contract.py b/examples/llm_server/python/tests/test_contract.py index a90415e276a..d5c6654ce80 100644 --- a/examples/llm_server/python/tests/test_contract.py +++ b/examples/llm_server/python/tests/test_contract.py @@ -185,11 +185,15 @@ def split_reasoning(text): @pytest.mark.parametrize("stream", [False, True]) +@pytest.mark.parametrize("has_extractor", [False, True]) @pytest.mark.parametrize("value", ["true", "false", 0, 1, 1.0, None, [], {}]) def test_return_reasoning_rejects_non_boolean_before_generation( - make_client, stream, value + make_client, stream, has_extractor, value ): - client, fake = make_client(max_named_sessions=1) + client, fake = make_client( + max_named_sessions=1, + reasoning_extractor=(lambda text: (None, text)) if has_extractor else None, + ) response = client.post( "/v1/chat/completions", json={ diff --git a/examples/llm_server/python/tests/test_server_launcher.py b/examples/llm_server/python/tests/test_server_launcher.py new file mode 100644 index 00000000000..1ba5da505b8 --- /dev/null +++ b/examples/llm_server/python/tests/test_server_launcher.py @@ -0,0 +1,62 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +import sys +from types import SimpleNamespace + +import pytest +import uvicorn + +from executorch.examples.llm_server.python import server + + +@pytest.mark.parametrize("configured", [False, True]) +def test_launcher_checks_assistant_header_before_starting_worker( + monkeypatch, caplog, configured +): + header = "<|start_header_id|>assistant<|end_header_id|>\n\n" + + class Tokenizer: + chat_template = "llama" + all_special_tokens = [] + + def apply_chat_template(self, messages, **kwargs): + return "<|start_header_id|>user<|end_header_id|>\n\n<|eot_id|>" + header + + monkeypatch.setitem( + sys.modules, + "transformers", + SimpleNamespace( + AutoTokenizer=SimpleNamespace(from_pretrained=lambda path: Tokenizer()) + ), + ) + argv = [ + "server", + "--worker-bin", + "worker", + "--model-path", + "model.pte", + "--tokenizer-path", + "tokenizer.json", + "--hf-tokenizer", + "local-tokenizer", + ] + if configured: + argv.extend(["--assistant-header", header]) + monkeypatch.setattr(sys, "argv", argv) + + def spawn(args): + warnings = [ + record for record in caplog.records if "Assistant header" in record.message + ] + assert len(warnings) == (0 if configured else 1) + return object() + + monkeypatch.setattr(server, "_spawn", spawn) + apps = [] + monkeypatch.setattr(uvicorn, "run", lambda app, **kwargs: apps.append(app)) + server.main() + assert len(apps) == 1 diff --git a/examples/llm_server/python/tests/test_sessions.py b/examples/llm_server/python/tests/test_sessions.py index f30e8731ac6..cb14875b730 100644 --- a/examples/llm_server/python/tests/test_sessions.py +++ b/examples/llm_server/python/tests/test_sessions.py @@ -307,7 +307,7 @@ def test_record_assistant_turn_replaces_stale_at_position(): generated_token_ids=[2], prior_turns=1, ) - assert [r["ids"] for r in t._turns["s"]] == [[1], [2]] + assert [r["ids"] for r in t._turns["s"].values()] == [[1], [2]] # regenerate turn 2 (same prior_turns) -> replaces stale [2], no stale tail t.record_assistant_turn( session_id="s", @@ -316,7 +316,7 @@ def test_record_assistant_turn_replaces_stale_at_position(): generated_token_ids=[3], prior_turns=1, ) - assert [r["ids"] for r in t._turns["s"]] == [[1], [3]] + assert [r["ids"] for r in t._turns["s"].values()] == [[1], [3]] def test_divergence_truncates_stale_tail(): @@ -357,7 +357,7 @@ def test_divergence_truncates_stale_tail(): template_kwargs=None, ) assert out.text == "X" # diverged -> plain text fallback - assert t._turns["s"] == [] # stale tail pruned from the first mismatch + assert t._turns["s"] == {} # stale tail pruned from the first mismatch class _HFToolSpecials: diff --git a/examples/llm_server/python/tests/test_template.py b/examples/llm_server/python/tests/test_template.py index fb25158342a..f9036600b4f 100644 --- a/examples/llm_server/python/tests/test_template.py +++ b/examples/llm_server/python/tests/test_template.py @@ -162,6 +162,25 @@ def test_fallback_ignores_kwargs_without_hf(): assert "<|im_start|>user" in out and out.endswith("<|im_start|>assistant\n") +@pytest.mark.parametrize("header", ["<|im_start|>assistant\n", ""]) +def test_generation_preamble_warns_once_for_missing_header(caplog, header): + template = ChatTemplate(allow_fallback=True, assistant_header=header) + template._hf = _FakeHF() + caplog.clear() + for kwargs in (None, None, {"enable_thinking": False}): + assert template.generation_preamble(kwargs) == "" + assert len(caplog.records) == 1 + assert "Assistant header" in caplog.text + assert "--assistant-header" in caplog.text + + +def test_generation_preamble_accepts_custom_header_without_warning(caplog): + template, _ = _template_with_gemma_tool_response_fake() + caplog.clear() + assert template.generation_preamble() == "<|channel>thought\n" + assert not caplog.records + + def test_tool_response_generation_prompt_disabled_by_default(): t, _ = _template_with_gemma_tool_response_fake(append=False) out = t.render([ChatMessage(role="tool", tool_call_id="c1", content="ok")]) diff --git a/examples/llm_server/python/tests/test_warm_resume_scaffold.py b/examples/llm_server/python/tests/test_warm_resume_scaffold.py index 4b91c45e4f1..95e1522ebe7 100644 --- a/examples/llm_server/python/tests/test_warm_resume_scaffold.py +++ b/examples/llm_server/python/tests/test_warm_resume_scaffold.py @@ -876,7 +876,13 @@ def test_plain_turn_splice_reproduces_resident_prefix(): "ORIGINAL", {"reasoning_content": "ORIGINAL "}, False, id="whitespace" ), pytest.param("ORIGINAL", {"reasoning_content": ""}, False, id="empty"), - pytest.param("ORIGINAL", {"reasoning_content": None}, False, id="null"), + pytest.param("ORIGINAL", {"reasoning_content": None}, True, id="null"), + pytest.param( + "line 1\nline 2", + {"reasoning_content": "line 1\r\nline 2"}, + False, + id="line-endings", + ), pytest.param(None, {"reasoning_content": None}, True, id="unchanged-null"), pytest.param(None, {"reasoning_content": ""}, False, id="null-to-empty"), ], @@ -1207,6 +1213,69 @@ def test_mistral_bracket_terminator_tail_falls_back(u1): assert enc(pi.text) == resident + enc("[INST]u2[/INST]") +@pytest.mark.parametrize( + "template_cls, header", + [(_FakeOtherHeader, _FakeOtherHeader.OHDR), (_FakeMistralNemo, "[/INST]")], + ids=["llama", "mistral"], +) +@pytest.mark.parametrize("configured", [False, True]) +def test_non_chatml_header_configuration_preserves_generated_bpe_tokens( + template_cls, header, configured +): + vocab = {chr(i): i for i in range(128)} + vocab["ab"] = len(vocab) + tokenizer = Tokenizer(models.BPE(vocab, merges=[("a", "b")])) + tokenizer.decoder = decoders.Fuse() + fake = template_cls() + + class TemplateTokenizer: + def apply_chat_template(self, messages, tools, **kwargs): + return fake.render([ChatMessage(**m) for m in messages], tools=tools) + + def enc(text): + return tokenizer.encode(text, add_special_tokens=False).ids + + template = ChatTemplate( + allow_fallback=True, assistant_header=header if configured else HDR + ) + template._hf = TemplateTokenizer() + state = OpenAITranscriptState(template) + first_prompt = template.render(_msgs(("user", "u1"))) + # Decoding preserves the answer, but re-encoding merges these two tokens. + generated = [vocab["a"], vocab["b"]] + assert tokenizer.decode(generated) == "ab" + assert generated != enc("ab") + resident = enc(first_prompt) + generated + state.record_assistant_turn( + session_id="s", + content="ab", + tool_calls=None, + generated_token_ids=generated, + prior_turns=0, + preamble=template.generation_preamble(), + ) + messages = _msgs(("user", "u1"), ("assistant", "ab"), ("user", "u2")) + rendered = template.render(messages) + prompt = state.build_prompt_input( + session_id="s", + messages=messages, + rendered_prompt=rendered, + tools=None, + template_kwargs=None, + ) + assembled = ( + enc(prompt.text) if prompt.text is not None else _assemble(prompt.segments, enc) + ) + assert tokenizer.decode(assembled) == rendered + if configured: + assert prompt.segments is not None + suffix = rendered[len(first_prompt + "ab") :] + assert assembled == resident + enc(suffix) + else: + assert prompt.text == rendered + assert assembled[: len(resident)] != resident + + # --- generation_preamble threads tools ------------------------------------ diff --git a/examples/llm_server/spec/README.md b/examples/llm_server/spec/README.md index 0f8bda86e64..8b4db6df3db 100644 --- a/examples/llm_server/spec/README.md +++ b/examples/llm_server/spec/README.md @@ -26,10 +26,16 @@ uses the worker's unset/random value. `model` must match the id returned by `/v1/models`; unknown ids return `404 model_not_found`. -For models with a reasoning extractor, `chat_template_kwargs.return_reasoning` -defaults to `true`. Set it to the boolean `false` to omit reasoning from the -response; this does not disable the model's reasoning computation. Non-boolean -values return `400 invalid_request_error` (`code: "invalid_value"`). +`chat_template_kwargs.return_reasoning` is an ExecuTorch response-visibility +control and defaults to `true`. For models with a reasoning extractor, set it +to `false` to omit reasoning from the response; the model still computes reasoning. +Unlike disabling reasoning separation in SGLang or llama.cpp, this suppresses +extracted reasoning instead of leaving it in `content`. + +The flag must be a JSON boolean for every model, including models without a +reasoning extractor. Non-boolean values return `400 invalid_request_error` +(`code: "invalid_value"`) before session admission or generation. This tightens +earlier behavior, which accepted non-boolean values as an opt-out. **Rejected** with `400 invalid_request_error` (`code: "unsupported_parameter"`) rather than silently ignored — a client relying on them would otherwise get @@ -95,9 +101,16 @@ suffix. All KV/resident state lives inside the worker/session, never the control plane. For named sessions, an unchanged assistant reply can reuse its original generated -token IDs. Clients may omit `reasoning_content` when echoing a reply. If the field -is supplied, it must match the value returned to that client; a changed value -(including explicit `null` or `""`) invalidates that turn's stored IDs and later -records. The updated history is then rendered normally. If an assistant boundary -cannot be verified, the server also uses the rendered text. The worker checks -the resulting token sequence before reusing KV state in either case. +token IDs. Clients may omit `reasoning_content` or send `null` when echoing a reply; +both allow replay of the original tokens, including reasoning. A supplied string +must exactly match the value returned to that client. String edits, including +whitespace or line-ending changes and `""`, invalidate that turn's stored IDs and +later records. Invalidated turns render normally and do not regain their old IDs +if the client restores the old history. Newly generated turns can be replayed at +their original assistant-turn indices. + +If an assistant boundary cannot be verified, the server uses the rendered text. +The generic launcher accepts `--assistant-header` and warns at startup if it is +absent from the template's generation prompt. The worker checks the resulting +token sequence before reusing KV state in either case; text fallback may still +reuse KV when the tokens match. diff --git a/examples/models/muse-glimmer/serving/serve.py b/examples/models/muse-glimmer/serving/serve.py index 7ce3122ecf3..068261f5123 100644 --- a/examples/models/muse-glimmer/serving/serve.py +++ b/examples/models/muse-glimmer/serving/serve.py @@ -84,9 +84,9 @@ def _extract_muse_glimmer_reasoning(text: str) -> tuple[str | None, str]: Nonempty thinking bodies retain their whitespace and are joined with a newline without adding channel framing. Interior text is not otherwise - sanitized. Warm resume compares any supplied reasoning against the returned + sanitized. Warm resume compares any supplied reasoning string against the returned text before splicing stored generated token ids, preserving the original - channel framing for unchanged echoes or omitted reasoning. + channel framing for unchanged echoes or omitted/null reasoning. """ matches = list(_MUSE_GLIMMER_ADDRESSED_HEADER_RE.finditer(text)) if not matches: diff --git a/examples/models/muse-glimmer/tests/test_serve.py b/examples/models/muse-glimmer/tests/test_serve.py index ef7e50a663f..dcc8a18ac22 100644 --- a/examples/models/muse-glimmer/tests/test_serve.py +++ b/examples/models/muse-glimmer/tests/test_serve.py @@ -705,7 +705,9 @@ async def _create_message(serving, request): @pytest.mark.parametrize("stream", [False, True]) -@pytest.mark.parametrize("echo_mode", ["edited", "unchanged", "omitted", "opt-out"]) +@pytest.mark.parametrize( + "echo_mode", ["edited", "unchanged", "omitted", "null", "opt-out", "opt-out-null"] +) def test_thinking_turn_replays_ids_only_for_unchanged_history(stream, echo_mode): # The worker reports non-terminal ids; the next rendered turn supplies EOS. # Two raw thinking blocks become one client-visible reasoning string. @@ -732,13 +734,15 @@ def test_thinking_turn_replays_ids_only_for_unchanged_history(stream, echo_mode) session_id="s", stream=stream, chat_template_kwargs=( - {"return_reasoning": False} if echo_mode == "opt-out" else None + {"return_reasoning": False} + if echo_mode.startswith("opt-out") + else None ), ), ) ) assert first["content"] == "156" - if echo_mode == "opt-out": + if echo_mode.startswith("opt-out"): assert "reasoning_content" not in first else: assert first["reasoning_content"] == "\nFirst step.\n\nSecond step." @@ -746,6 +750,8 @@ def test_thinking_turn_replays_ids_only_for_unchanged_history(stream, echo_mode) first["reasoning_content"] = "Use this corrected reasoning instead." elif echo_mode == "omitted": first.pop("reasoning_content") + elif echo_mode in ("null", "opt-out-null"): + first["reasoning_content"] = None request2 = ChatCompletionRequest( messages=[u1, ChatMessage(**first), ChatMessage(role="user", content="And?")], session_id="s", @@ -762,6 +768,54 @@ def test_thinking_turn_replays_ids_only_for_unchanged_history(stream, echo_mode) assert _assemble(prompt2) == resident + _benc(trailing) +@pytest.mark.parametrize("stream", [False, True]) +@pytest.mark.parametrize("edited_field", ["content", "reasoning_content"]) +@pytest.mark.parametrize("restore_history", [False, True]) +def test_replay_recovers_after_edit_without_restoring_stale_turns( + stream, edited_field, restore_history +): + raw = [ + f" to=self<|message|>reasoning {i}<|eom|>" + f"<|start|>assistant to=user<|message|>a{i}" + for i in range(1, 6) + ] + generated = [_benc(text) for text in raw] + template = _HarmonyTemplate() + serving, runtime = _glimmer_serving( + [{"text": text, "gen_ids": ids} for text, ids in zip(raw, generated)], + template=template, + ) + + async def conversation(): + messages = [] + for i in range(5): + if i == 3: + original = messages[3] + messages[3] = original.model_copy(update={edited_field: "EDITED"}) + elif i == 4 and restore_history: + messages[3] = original + messages.append(ChatMessage(role="user", content=f"u{i + 1}")) + response = await _create_message( + serving, + ChatCompletionRequest( + messages=list(messages), session_id="s", stream=stream + ), + ) + assert response["content"] == f"a{i + 1}" + prompt = runtime.prompts[-1] + assert _assemble(prompt) == _benc(template.render(messages)) + expected = generated[:i] if i < 3 else [generated[0]] + if i == 4: + expected.append(generated[3]) + if not restore_history: + resident = _assemble(runtime.prompts[3]) + generated[3] + assert _assemble(prompt)[: len(resident)] == resident + assert [s["ids"] for s in prompt.segments or [] if "ids" in s] == expected + messages.append(ChatMessage(**response)) + + asyncio.run(conversation()) + + def test_extract_muse_glimmer_reasoning_plain_text_fallback(): text = "The capital of France is Paris." assert serve._extract_muse_glimmer_reasoning(text) == (None, text)