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..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 @@ -114,7 +122,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 +134,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/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 a6bf5b75e6d..660efbc0920 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 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. """ import hashlib @@ -111,9 +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 - # session_id -> [{"fp": str, "ids": list[int] | None}, ...] (one per - # assistant turn we produced, in order). 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: @@ -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,40 +270,37 @@ 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 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] - if self._assistant_fingerprint(m.content, m.tool_calls) != stored[k]["fp"]: - diverged_at = k # this stored turn and every later one are stale + if self._assistant_fingerprint(m.content, m.tool_calls) != record["fp"] or ( + m.reasoning_content is not None + and self._reasoning_fingerprint(m.reasoning_content) + != record["reasoning_fp"] + ): + # 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 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 - # 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 = { @@ -351,25 +352,28 @@ 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, []) - del turns[prior_turns:] - turns.append( - { - "fp": self._assistant_fingerprint(content, tool_calls), - "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 0530719e922..bf2d1da17bb 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/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 61cb671fee3..319220ef797 100644 --- a/examples/llm_server/python/serving_chat.py +++ b/examples/llm_server/python/serving_chat.py @@ -139,12 +139,9 @@ 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 {} - 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 +349,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 +558,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 +749,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..d5c6654ce80 100644 --- a/examples/llm_server/python/tests/test_contract.py +++ b/examples/llm_server/python/tests/test_contract.py @@ -134,6 +134,86 @@ 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("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, has_extractor, value +): + 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={ + "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_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_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..95e1522ebe7 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,194 @@ 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}, 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"), + ], +) +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 +1063,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 +1137,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 +1210,70 @@ 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]") + + +@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 ------------------------------------ @@ -1016,14 +1309,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 +1361,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..8b4db6df3db 100644 --- a/examples/llm_server/spec/README.md +++ b/examples/llm_server/spec/README.md @@ -26,6 +26,17 @@ uses the worker's unset/random value. `model` must match the id returned by `/v1/models`; unknown ids return `404 model_not_found`. +`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 wrong behavior: `n` (> 1), `reasoning_effort`, @@ -41,6 +52,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 +63,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 +99,18 @@ 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` 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/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..068261f5123 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 string against the returned + text before splicing stored generated token ids, preserving the original + channel framing for unchanged echoes or omitted/null 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..dcc8a18ac22 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,137 @@ 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", "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. 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.startswith("opt-out") + else None + ), + ), ) ) + assert first["content"] == "156" + if echo_mode.startswith("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") + 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", + ) + 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) + + +@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():