From dda11089a6549fb50421e3743f3a489e81bb709a Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 19:07:37 +0000 Subject: [PATCH 01/16] Preserve literal assistant content in Qwen thinking templates --- docs/features/additional-histories.mdx | 14 +- src/art_inference/chat_template.py | 27 ++- tests/unit/test_literal_reasoning_content.py | 211 ++++++++++++++++++ .../trajectories/test_literal_thinking_off.py | 13 +- 4 files changed, 259 insertions(+), 6 deletions(-) create mode 100644 tests/unit/test_literal_reasoning_content.py diff --git a/docs/features/additional-histories.mdx b/docs/features/additional-histories.mdx index 24802ace7..5e8deeeba 100644 --- a/docs/features/additional-histories.mdx +++ b/docs/features/additional-histories.mdx @@ -33,6 +33,16 @@ with `chat_template_kwargs={"preserve_thinking": False}`. Additional histories remain useful for custom or externally managed templates that do not expose a prior-thinking preservation option. +For supported Qwen templates, reasoning belongs in the structured +`reasoning_content` field. Assistant `content` remains literal, including +`` and `` anywhere in that content; ART does not infer reasoning +from those strings. This also applies when the next response has thinking +enabled or when prior reasoning is explicitly omitted. A legacy adapter that +knows its response uses a leading reasoning envelope should split that known +format into `reasoning_content` and `content` before rendering. A leading tag +pair alone cannot establish that format. Structured reasoning and the template's +generation-prompt defaults retain their existing behavior. + By splitting each turn into a separate history, you can preserve these tokens for training: ```python @@ -44,7 +54,7 @@ trajectory = Trajectory( messages_and_choices=[ # First turn with thinking {"role": "user", "content": "What is 2+2?"}, - {"role": "assistant", "content": "I need to add 2 and 24"} + {"role": "assistant", "reasoning_content": "I need to add 2 and 2", "content": "4"} ], additional_histories=[ LegacyHistory( @@ -53,7 +63,7 @@ trajectory = Trajectory( {"role": "user", "content": "What is 2+2?"}, {"role": "assistant", "content": "4"}, {"role": "user", "content": "What is 3+3?"}, - {"role": "assistant", "content": "I need to add 3 and 36"} + {"role": "assistant", "reasoning_content": "I need to add 3 and 3", "content": "6"} ] ) ] diff --git a/src/art_inference/chat_template.py b/src/art_inference/chat_template.py index 20dd21c7e..bec0d5795 100644 --- a/src/art_inference/chat_template.py +++ b/src/art_inference/chat_template.py @@ -25,10 +25,24 @@ "reasoning_content and ((preserve_thinking is defined and preserve_thinking is " "true) or loop.index0 > ns.last_user_index)" ) +# These operations infer reasoning from arbitrary assistant content and can +# discard everything before the last or between repeated tags. +# Match the operations, not a model revision or the text of a particular answer. +_QWEN_INLINE_REASONING = re.compile( + r"\s*".join( + r"\{%[-+]?\s*" + re.escape(statement) + r"\s*[-+]?%\}" + for statement in ( + "if '' in content", + "set reasoning_content = content.split('')[0].rstrip('\\n').split('')[-1].lstrip('\\n')", + "set content = content.split('')[-1].lstrip('\\n')", + "endif", + ) + ) +) def chat_template_with_preserved_thinking(chat_template: object) -> object: - """Preserve prior reasoning by default, while respecting explicit opt-outs.""" + """Preserve structured reasoning without interpreting tags in plain content.""" if isinstance(chat_template, dict): return { name: chat_template_with_preserved_thinking(template) @@ -36,6 +50,17 @@ def chat_template_with_preserved_thinking(chat_template: object) -> object: } if not isinstance(chat_template, str): return chat_template + chat_template, inline_parsers = _QWEN_INLINE_REASONING.subn("", chat_template) + if inline_parsers: + # Disabling reasoning preservation may omit a structured reasoning + # field, but must not trim the visible assistant answer. + chat_template = chat_template.replace( + "if preserve_thinking and message.role == 'assistant'", + "if message.role == 'assistant'", + ).replace( + "set content = render_content(message.content, true)|trim", + "set content = (render_content(message.content, true) if message.role == 'assistant' else render_content(message.content, true)|trim)", + ) replacements = ( ( _QWEN_DROP_PRIOR_THINKING, diff --git a/tests/unit/test_literal_reasoning_content.py b/tests/unit/test_literal_reasoning_content.py new file mode 100644 index 000000000..7f79e7dab --- /dev/null +++ b/tests/unit/test_literal_reasoning_content.py @@ -0,0 +1,211 @@ +from copy import deepcopy +import hashlib +from pathlib import Path + +from jinja2.sandbox import ImmutableSandboxedEnvironment +import pytest + +from art_inference.chat_template import ( + chat_template_with_preserved_thinking, + default_chat_template_kwargs_for_template, +) + +_TEMPLATE = ( + Path(__file__).parents[1] / "fixtures/qwen35_preserved_thinking.jinja" +).read_text() +_FIXED = chat_template_with_preserved_thinking(_TEMPLATE) +_USER = {"role": "user", "content": "A public question."} +_LITERALS = ( + "plain answer", + "thoughtanswer", + "prefixliteralsuffix", + "answerliteral", + "prefixmiddlesuffix", + "onetwo", + "nested", + "unclosed", + "unopened", + "", + "", + "\n before café 漢字 🦉 after \n", + "", +) + + +def _render(template, messages, **kwargs): + def refuse(message): + raise ValueError(message) + + env = ImmutableSandboxedEnvironment( + trim_blocks=True, lstrip_blocks=True, extensions=["jinja2.ext.loopcontrols"] + ) + return env.from_string(template).render( + messages=messages, raise_exception=refuse, **kwargs + ) + + +@pytest.mark.parametrize("thinking", [False, True]) +@pytest.mark.parametrize("preserve", [False, True]) +@pytest.mark.parametrize("content", _LITERALS) +def test_plain_content_is_literal_in_every_mode(content, thinking, preserve): + messages = [_USER, {"role": "assistant", "content": content}] + before = deepcopy(messages) + rendered = _render( + _FIXED, messages, enable_thinking=thinking, preserve_thinking=preserve + ) + assert rendered.endswith(content + "<|im_end|>\n") + assert messages == before + # Changing the next turn's thinking mode never reinterprets history. + assert rendered == _render( + _FIXED, messages, enable_thinking=not thinking, preserve_thinking=preserve + ) + + +@pytest.mark.parametrize("thinking", [False, True]) +@pytest.mark.parametrize("preserve", [False, True]) +@pytest.mark.parametrize( + "reasoning", [None, "", "reasoned\n", "literal reasoning text\n"] +) +def test_structured_reasoning_and_explicit_empty_field_keep_existing_behavior( + thinking, preserve, reasoning +): + messages = [ + _USER, + {"role": "assistant", "content": "answer", "reasoning_content": reasoning}, + {"role": "user", "content": "next"}, + ] + kwargs = dict(enable_thinking=thinking, preserve_thinking=preserve) + before = deepcopy(messages) + assert _render(_FIXED, messages, **kwargs) == _render(_TEMPLATE, messages, **kwargs) + assert messages == before + rendered = _render(_FIXED, messages, **kwargs) + if reasoning: + assert (reasoning in rendered) == preserve + + +@pytest.mark.parametrize("thinking", [False, True]) +@pytest.mark.parametrize("preserve", [False, True]) +def test_proven_legacy_encoding_uses_existing_structured_fields(thinking, preserve): + # This fixture declares the old encoding. The renderer cannot infer that + # declaration from an indistinguishable literal string in plain content. + legacy = {"role": "assistant", "content": "\nthought\n\n\nanswer"} + structured = { + "role": "assistant", + "reasoning_content": "thought\n", + "content": "answer", + } + kwargs = dict(enable_thinking=thinking, preserve_thinking=preserve) + later = {"role": "user", "content": "next"} + assert _render(_FIXED, [_USER, structured, later], **kwargs) == _render( + _TEMPLATE, [_USER, legacy, later], **kwargs + ) + assert legacy["content"] in _render(_FIXED, [_USER, legacy], **kwargs) + + +@pytest.mark.parametrize("thinking", [False, True]) +@pytest.mark.parametrize("preserve", [False, True]) +@pytest.mark.parametrize("content", [None, "", "beforeliteralafter"]) +def test_tool_call_and_continuation_keep_content_and_arguments( + thinking, preserve, content +): + assistant = { + "role": "assistant", + "content": content, + "tool_calls": [ + {"function": {"name": "lookup", "arguments": {"q": ""}}} + ], + } + messages = [_USER, assistant] + before = deepcopy(messages) + kwargs = dict(enable_thinking=thinking, preserve_thinking=preserve) + rendered = _render(_FIXED, messages, **kwargs) + assert "" in rendered + assert "\n\n" in rendered + if content: + assert content in rendered + continued = _render( + _FIXED, + [ + *messages, + {"role": "tool", "content": "result"}, + {"role": "assistant", "content": "nextliteralanswer"}, + ], + **kwargs, + ) + assert continued.startswith(rendered) + assert messages == before + + +@pytest.mark.parametrize("thinking", [False, True]) +@pytest.mark.parametrize("preserve", [False, True]) +def test_generation_prompt_and_preserved_history_prefix_are_stable(thinking, preserve): + kwargs = dict( + enable_thinking=thinking, preserve_thinking=preserve, add_generation_prompt=True + ) + assert _render(_FIXED, [_USER], **kwargs) == _render(_TEMPLATE, [_USER], **kwargs) + messages = [ + _USER, + {"role": "assistant", "content": "plainliteraltail"}, + ] + completed = _render(_FIXED, messages, preserve_thinking=preserve) + continuation = _render(_FIXED, [*messages, _USER], **kwargs) + if preserve: + assert continuation.startswith(completed) + else: + # Explicit opt-out still removes the previous turn's reasoning scaffold; + # the visible body is unchanged, not a newly promised full-token prefix. + assert messages[-1]["content"] + "<|im_end|>\n" in continuation + assert not continuation.startswith(completed) + + +def test_actual_template_operation_and_public_e2ac_shaped_regression(): + assert ( + hashlib.sha256(_TEMPLATE.encode()).hexdigest() + == "098047d425a6673b1fe1a82a197a481616e53a283beaa8cb76cbb74d38ca6644" + ) + # Public text with the captured branch's shape; no private text or IDs. + prefix = "P" * 4149 + body = prefix + "\n" + "R" * 747 + "\n\n\n" + "A" * 1161 + messages = [_USER, {"role": "assistant", "content": body}] + original = _render( + _TEMPLATE, messages, enable_thinking=False, preserve_thinking=True + ) + assert prefix not in original + assert body in _render( + _FIXED, messages, enable_thinking=False, preserve_thinking=True + ) + + +def test_configuration_is_idempotent_and_does_not_change_defaults_or_other_templates(): + assert _FIXED != _TEMPLATE + assert chat_template_with_preserved_thinking(_FIXED) == _FIXED + assert default_chat_template_kwargs_for_template( + _FIXED + ) == default_chat_template_kwargs_for_template(_TEMPLATE) + other = "{% for message in messages %}{{ message.content }}{% endfor %}" + assert chat_template_with_preserved_thinking(other) == other + assert chat_template_with_preserved_thinking( + {"default": _TEMPLATE, "other": other} + ) == {"default": _FIXED, "other": other} + + +def test_unconfigured_template_receives_the_same_correction(): + # Reverse the prior preservation-only rewrite of this public fixture. + raw = ( + _TEMPLATE.replace( + "{%- set preserve_thinking = preserve_thinking | default(true) -%}", "" + ) + .replace( + "(render_content(message.content, true) if preserve_thinking and message.role == 'assistant' else render_content(message.content, true)|trim)", + "render_content(message.content, true)|trim", + ) + .replace( + "{%- if not preserve_thinking or message.reasoning_content is not string %}{%- set reasoning_content = reasoning_content|trim %}{%- endif %}", + "{%- set reasoning_content = reasoning_content|trim %}", + ) + .replace( + "('\\n\\n' if preserve_thinking and message.reasoning_content is string and reasoning_content else '\\n\\n\\n')", + "'\\n\\n\\n'", + ) + ) + assert chat_template_with_preserved_thinking(raw) == _FIXED diff --git a/tests/unit/trajectories/test_literal_thinking_off.py b/tests/unit/trajectories/test_literal_thinking_off.py index 16188c733..03e567e21 100644 --- a/tests/unit/trajectories/test_literal_thinking_off.py +++ b/tests/unit/trajectories/test_literal_thinking_off.py @@ -148,6 +148,9 @@ def test_native_thinking_off_retains_literal_content( patch.setattr( _tokenize, "_preserve_literal_thinking_off_content", lambda *args: None ) + patch.setattr( + _tokenize, "chat_template_with_preserved_thinking", lambda value: value + ) _outcome(history, tokenizer) assert content not in tokenizer.rendered[0] tokenizer.calls.clear() @@ -162,7 +165,7 @@ def test_native_thinking_off_retains_literal_content( required = tr.TokenFlag.EXACT | tr.TokenFlag.ASSISTANT | tr.TokenFlag.OUTPUT assert all(tokenized.flags[i] & required == required for i in sampled) assert not any(flag & tr.TokenFlag.STOP for flag in tokenized.flags) - assert tokenizer.calls[0][-1]["reasoning_content"] == "" + assert tokenizer.calls[0][-1].get("reasoning_content", "") == "" assert tokenizer.calls[0][-1]["content"] == content assert history.model_dump(mode="python") == original @@ -249,7 +252,7 @@ def test_explicit_empty_reasoning_is_preserved(field: str) -> None: ) == _LITERAL ) - assert tokenizer.calls[0][-1]["reasoning_content"] == "" + assert tokenizer.calls[0][-1].get("reasoning_content", "") == "" assert history.model_dump(mode="python") == original @@ -278,7 +281,8 @@ def test_mixed_history_uses_each_generations_own_request( _outcome(history, tokenizer) rendered_messages = tokenizer.calls[0] assert "reasoning_content" not in rendered_messages[1] - assert rendered_messages[3]["reasoning_content"] == "" + assert "reasoning_content" not in rendered_messages[3] + assert _LITERAL in tokenizer.rendered[0] assert [message["content"] for message in rendered_messages] == [ message["content"] for message in history.messages ] @@ -405,6 +409,9 @@ def observe(*args: Any, **kwargs: Any) -> tr.TokenizedHistory | None: patch.setattr( _tokenize, "_preserve_literal_thinking_off_content", lambda *args: None ) + patch.setattr( + _tokenize, "chat_template_with_preserved_thinking", lambda value: value + ) _outcome(history, tokenizer) boundary, old_exact = observed[0] stored = list(boundary.tail + boundary.following) From ddab2360d74172c1d56b5e25cd586555b0b15d7e Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 19:09:49 +0000 Subject: [PATCH 02/16] Constrain literal-thinking rewrite to its supported template syntax --- src/art_inference/chat_template.py | 21 ++++++++++++------- tests/unit/test_literal_reasoning_content.py | 22 ++++++++++++++++++++ 2 files changed, 35 insertions(+), 8 deletions(-) diff --git a/src/art_inference/chat_template.py b/src/art_inference/chat_template.py index bec0d5795..1777c3978 100644 --- a/src/art_inference/chat_template.py +++ b/src/art_inference/chat_template.py @@ -50,17 +50,22 @@ def chat_template_with_preserved_thinking(chat_template: object) -> object: } if not isinstance(chat_template, str): return chat_template - chat_template, inline_parsers = _QWEN_INLINE_REASONING.subn("", chat_template) + # This source rewrite is deliberately conservative, not a Jinja parser. + # In raw/comment-containing templates the same text might be literal data. + inline_parsers = 0 + if not re.search(r"\{#|\{%[-+]?\s*raw\b", chat_template): + chat_template, inline_parsers = _QWEN_INLINE_REASONING.subn("", chat_template) if inline_parsers: # Disabling reasoning preservation may omit a structured reasoning # field, but must not trim the visible assistant answer. - chat_template = chat_template.replace( - "if preserve_thinking and message.role == 'assistant'", - "if message.role == 'assistant'", - ).replace( - "set content = render_content(message.content, true)|trim", - "set content = (render_content(message.content, true) if message.role == 'assistant' else render_content(message.content, true)|trim)", - ) + for content in ( + "render_content(message.content, true)|trim", + "(render_content(message.content, true) if preserve_thinking and message.role == 'assistant' else render_content(message.content, true)|trim)", + ): + chat_template = chat_template.replace( + "{%- set content = " + content + " %}", + "{%- set content = (render_content(message.content, true) if message.role == 'assistant' else render_content(message.content, true)|trim) %}", + ) replacements = ( ( _QWEN_DROP_PRIOR_THINKING, diff --git a/tests/unit/test_literal_reasoning_content.py b/tests/unit/test_literal_reasoning_content.py index 7f79e7dab..d39981d23 100644 --- a/tests/unit/test_literal_reasoning_content.py +++ b/tests/unit/test_literal_reasoning_content.py @@ -209,3 +209,25 @@ def test_unconfigured_template_receives_the_same_correction(): ) ) assert chat_template_with_preserved_thinking(raw) == _FIXED + + +@pytest.mark.parametrize("wrapper", [("{% raw %}", "{% endraw %}"), ("{#", "#}")]) +def test_inline_operation_as_raw_or_comment_text_is_not_rewritten(wrapper): + from art_inference.chat_template import _QWEN_INLINE_REASONING + + operation = _QWEN_INLINE_REASONING.search(_TEMPLATE).group() + template = wrapper[0] + operation + wrapper[1] + assert chat_template_with_preserved_thinking(template) == template + + +def test_other_structured_reasoning_condition_is_not_rewritten(): + gate = "{% if preserve_thinking and message.role == 'assistant' %}{{ message.reasoning_content }}{% endif %}" + template = _TEMPLATE + "{% for message in messages %}" + gate + "{% endfor %}" + fixed = chat_template_with_preserved_thinking(template) + assert gate in fixed + messages = [ + _USER, + {"role": "assistant", "content": "answer", "reasoning_content": "prior reason"}, + _USER, + ] + assert "prior reason" not in _render(fixed, messages, preserve_thinking=False) From 52cb457a21c1b09818ff92ee0418950afae2f88b Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 19:14:07 +0000 Subject: [PATCH 03/16] Bind thinking-template edits to executable Jinja blocks --- src/art_inference/chat_template.py | 58 ++++++++++++++------ tests/unit/test_literal_reasoning_content.py | 27 +++++++++ 2 files changed, 69 insertions(+), 16 deletions(-) diff --git a/src/art_inference/chat_template.py b/src/art_inference/chat_template.py index 1777c3978..fcf8a1d0d 100644 --- a/src/art_inference/chat_template.py +++ b/src/art_inference/chat_template.py @@ -41,6 +41,47 @@ ) +def _without_inline_reasoning_parser(template: str) -> str: + matches = list(_QWEN_INLINE_REASONING.finditer(template)) + if not matches: + return template + from jinja2 import Environment, TemplateSyntaxError + + # Only executable block tokens may be edited. The same spelling inside a + # quoted expression, raw block or comment is literal template data. + starts: set[int] = set() + cursor = 0 + try: + for _, kind, value in Environment().lex(template): + start = template.find(value, cursor) + if start < 0 or template[cursor:start].strip(): + return template # Lexer normalization could not be source-joined. + if kind == "block_begin": + starts.add(start) + cursor = start + len(value) + except TemplateSyntaxError: + return template # Leave invalid templates to their existing renderer. + edits = { + (match.start(), match.end()): "" for match in matches if match.start() in starts + } + if not edits: + return template + # Dropping structured reasoning must not trim the visible assistant body. + for content in ( + "render_content(message.content, true)|trim", + "(render_content(message.content, true) if preserve_thinking and message.role == 'assistant' else render_content(message.content, true)|trim)", + ): + statement = "{%- set content = " + content + " %}" + for match in re.finditer(re.escape(statement), template): + if match.start() in starts: + edits[match.span()] = ( + "{%- set content = (render_content(message.content, true) if message.role == 'assistant' else render_content(message.content, true)|trim) %}" + ) + for (start, end), replacement in sorted(edits.items(), reverse=True): + template = template[:start] + replacement + template[end:] + return template + + def chat_template_with_preserved_thinking(chat_template: object) -> object: """Preserve structured reasoning without interpreting tags in plain content.""" if isinstance(chat_template, dict): @@ -50,22 +91,7 @@ def chat_template_with_preserved_thinking(chat_template: object) -> object: } if not isinstance(chat_template, str): return chat_template - # This source rewrite is deliberately conservative, not a Jinja parser. - # In raw/comment-containing templates the same text might be literal data. - inline_parsers = 0 - if not re.search(r"\{#|\{%[-+]?\s*raw\b", chat_template): - chat_template, inline_parsers = _QWEN_INLINE_REASONING.subn("", chat_template) - if inline_parsers: - # Disabling reasoning preservation may omit a structured reasoning - # field, but must not trim the visible assistant answer. - for content in ( - "render_content(message.content, true)|trim", - "(render_content(message.content, true) if preserve_thinking and message.role == 'assistant' else render_content(message.content, true)|trim)", - ): - chat_template = chat_template.replace( - "{%- set content = " + content + " %}", - "{%- set content = (render_content(message.content, true) if message.role == 'assistant' else render_content(message.content, true)|trim) %}", - ) + chat_template = _without_inline_reasoning_parser(chat_template) replacements = ( ( _QWEN_DROP_PRIOR_THINKING, diff --git a/tests/unit/test_literal_reasoning_content.py b/tests/unit/test_literal_reasoning_content.py index d39981d23..0df6826a9 100644 --- a/tests/unit/test_literal_reasoning_content.py +++ b/tests/unit/test_literal_reasoning_content.py @@ -231,3 +231,30 @@ def test_other_structured_reasoning_condition_is_not_rewritten(): _USER, ] assert "prior reason" not in _render(fixed, messages, preserve_thinking=False) + + +def test_inline_operation_inside_quoted_expression_is_literal(): + from art_inference.chat_template import _QWEN_INLINE_REASONING + + operation = _QWEN_INLINE_REASONING.search(_TEMPLATE).group().replace("\n", " ") + template = '{{ "' + operation + '" }}' + fixed = chat_template_with_preserved_thinking(template) + assert fixed == template + assert "set reasoning_content = content.split" in _render(fixed, []) + + +@pytest.mark.parametrize( + "wrapper", [("{% raw %}", "{% endraw %}"), ("{#", "#}"), ('{{ "', '" }}')] +) +def test_mixed_executable_and_literal_operations_only_changes_executable(wrapper): + from art_inference.chat_template import _QWEN_INLINE_REASONING + + operation = _QWEN_INLINE_REASONING.search(_TEMPLATE).group().replace("\n", " ") + literal = wrapper[0] + operation + wrapper[1] + template = _TEMPLATE + literal + fixed = chat_template_with_preserved_thinking(template) + assert fixed == _FIXED + literal + assert "headliteraltail" in _render( + fixed, + [_USER, {"role": "assistant", "content": "headliteraltail"}], + ) From 225bb8ac24895c9ffd9e112a5e0c6e24412e349f Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 19:37:37 +0000 Subject: [PATCH 04/16] Use recorded boundaries and preserve sampled conditioning by default --- docs/features/additional-histories.mdx | 26 + src/art/trajectories/_tokenize.py | 598 +++++++++++++++--- tests/unit/test_literal_reasoning_content.py | 13 +- .../trajectories/test_literal_thinking_off.py | 32 +- .../trajectories/test_recorded_boundaries.py | 555 ++++++++++++++++ tests/unit/trajectories/test_tokenize.py | 40 +- 6 files changed, 1144 insertions(+), 120 deletions(-) create mode 100644 tests/unit/trajectories/test_recorded_boundaries.py diff --git a/docs/features/additional-histories.mdx b/docs/features/additional-histories.mdx index 5e8deeeba..fd20373f4 100644 --- a/docs/features/additional-histories.mdx +++ b/docs/features/additional-histories.mdx @@ -144,6 +144,32 @@ trajectory = Trajectory( ) ``` +## Recorded exchange histories + +`art.tokenize` and `trajectory.tokenize` use recorded prompt and response token IDs +for unchanged, complete exchange histories. Recorded logprobs belong to those +exact conditioned tokens. Chat, Responses, Messages, and Completions keep their +existing protocol-specific projection rules; no separate tokenization API or +native-representation option is needed. `multi_history=True` preserves the +histories selected by the trajectory, including their order and model selection. + +Templates still own unrecorded separators, role masks, and synthetic stop tokens. +For supported Chat boundaries, ART decodes the recorded body and encodes only the +unrecorded separator instead of re-tokenizing the whole conversation. It checks +that the separator reproduces the next recorded prompt exactly. Edited contexts, +explicit template overrides, incomplete projections, and unsupported templates +continue through the generic rendering path and its source validation. + +A response copied into a later, shortened prompt is output provenance, but it is +not a fresh sample under that new prompt. ART keeps its `OUTPUT`, `ASSISTANT`, +`EXACT`, and proven `STOP` flags while removing `SAMPLED` and the old conditional +logprob. This requires the complete original sampled occurrence to remain in an +earlier selected history; otherwise unchanged native replay is refused. The +original occurrence retains its logprobs and ownership, including recorded NaNs +before finite-value filtering. Tokenizing only the shortened view cannot prove +that coverage; tokenize the containing trajectory instead. This correction does +not change how generic output/SFT masks include copied assistant content. + ## How It Works ### Tokenization Process diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index bfdebad26..e13db82be 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -3588,6 +3588,7 @@ def _tokenize_exact_responses_history( base_model: str | None, tokenizer: Tokenizer | None, _trace: _TraceBuilder | None = None, + _prior: Sequence[tuple[TokenizedHistory, _HistoryTokenizationTrace]] = (), ) -> TokenizedHistory | None: generation_keys: list[tuple[ResponsesExchange, int]] = [] retained_output_indices: dict[tuple[int, int], set[int]] = {} @@ -3619,8 +3620,29 @@ def _tokenize_exact_responses_history( output = generation.output_token_ids if prompt is None or output is None: return None + source = next( + item + for item in history.input_sources + if item is not None + and item.exchange is exchange + and item.generation_index == generation_index + ) + context_only = False retained = retained_output_indices.get((id(exchange), generation_index), set()) - if retained != set(generation.output_indices): + following_prompt = None + if position + 1 < len(generation_keys): + following_exchange, following_index = generation_keys[position + 1] + following_prompt = _response_generations(following_exchange.response)[ + following_index + ].prompt_token_ids + from ._history import _retains_output_suffix + + copied_suffix = ( + following_prompt is not None + and following_prompt[: len(prompt) + len(output)] != [*prompt, *output] + and _retains_output_suffix(prompt, output, following_prompt) + ) + if retained != set(generation.output_indices) or copied_suffix: if position + 1 >= len(generation_keys): return None next_exchange, next_generation_index = generation_keys[position + 1] @@ -3638,7 +3660,16 @@ def _tokenize_exact_responses_history( ) if retained_suffix is None: return None + context_only = retained_suffix[0] != output + if context_only and not _complete_source_is_represented( + source, prompt, output, generation.output_logprobs, _prior + ): + raise ValueError( + "A copied Responses suffix requires its complete original sampled occurrence in the selected trajectory" + ) output, output_logprobs = retained_suffix + if context_only: + output_logprobs = [math.nan] * len(output) output_text = None else: output_logprobs = generation.output_logprobs @@ -3682,7 +3713,7 @@ def _tokenize_exact_responses_history( flags.extend( [ TokenFlag.EXACT - | TokenFlag.SAMPLED + | (TokenFlag(0) if context_only else TokenFlag.SAMPLED) | TokenFlag.ASSISTANT | TokenFlag.OUTPUT ] @@ -3701,15 +3732,28 @@ def _tokenize_exact_responses_history( if source is None: raise AssertionError("Responses generation has no history source") source_key = _sampled_source_key(source) - source_keys.extend([source_key] * len(output)) + source_keys.extend([None if context_only else source_key] * len(output)) sources[source_key] = source - sampled_outputs.append( - _SampledOutput( - text=output_text, - token_ids=list(output), - start=len(token_ids) - len(output), + if context_only: + stop_count = _sampled_stop_suffix( + generation.output_token_ids or [], + source=source, + source_key=source_key, + tokenizer=tokenizer, + ) + for offset in range( + max(len(token_ids) - len(output), len(token_ids) - stop_count), + len(token_ids), + ): + flags[offset] |= TokenFlag.STOP + else: + sampled_outputs.append( + _SampledOutput( + text=output_text, + token_ids=list(output), + start=len(token_ids) - len(output), + ) ) - ) _mark_sampled_stops( token_ids, flags, @@ -3851,6 +3895,15 @@ def _chat_source_prompt_tokens(source: object) -> list[int] | None: return None +def _chat_source_record( + source: object, +) -> tuple[list[int] | None, list[int] | None, list[float]]: + exchange = getattr(source, "exchange", None) + if isinstance(exchange, ChatCompletionsExchange): + return _chat_choice_tokens(_chat_choice(source), exchange.response) + return _chat_source_prompt_tokens(source), *_chat_source_full_tokens(source) + + def _source_is_sampled(source: object) -> bool: exchange = getattr(source, "exchange", None) if isinstance(exchange, ChatCompletionsExchange): @@ -4118,6 +4171,151 @@ def _next_assistant_span_start( ) +def _complete_source_is_represented( + source: object, + prompt: list[int], + output: list[int], + logprobs: list[float], + prior: Sequence[tuple[TokenizedHistory, _HistoryTokenizationTrace]], +) -> bool: + """Prove ownership of the original edge before treating a copy as context.""" + key = _sampled_source_key(source) + exchange = getattr(source, "exchange", None) + required = ( + TokenFlag.EXACT | TokenFlag.SAMPLED | TokenFlag.ASSISTANT | TokenFlag.OUTPUT + ) + expected_lp = logprobs if len(logprobs) == len(output) else [math.nan] * len(output) + end = len(prompt) + len(output) + for previous, trace in prior: + owner = trace.sources.get(key) + if ( + previous.model != getattr(exchange, "model", None) + or getattr(owner, "exchange", None) is not exchange + or getattr(owner, "choice_index", None) + != getattr(source, "choice_index", None) + or trace.source_keys[len(prompt) : end] != [key] * len(output) + or previous.tokens[:end] != [*prompt, *output] + or any( + flag & required != required + for flag in previous.flags[len(prompt) : end] + ) + ): + continue + if all( + left == right or math.isnan(left) and math.isnan(right) + for left, right in zip( + previous.logprobs[len(prompt) : end], expected_lp, strict=True + ) + ): + return True + return False + + +def _source_native_record( + source: object, +) -> tuple[list[int] | None, list[int] | None, list[float]]: + exchange = getattr(source, "exchange", None) + if isinstance(exchange, ResponsesExchange): + index = getattr(source, "generation_index", None) + generations = _response_generations(exchange.response) + if isinstance(index, int) and 0 <= index < len(generations): + generation = generations[index] + return ( + generation.prompt_token_ids, + generation.output_token_ids, + generation.output_logprobs, + ) + return _chat_source_record(source) + + +def _source_native_prefix(source: object) -> tuple[list[int] | None, list[int] | None]: + # A capability preflight only: the normal source readers still validate the + # records before assembly. Avoid decoding LP carriers just to detect copies. + exchange = getattr(source, "exchange", None) + if isinstance(exchange, ChatCompletionsExchange): + choice = _chat_choice(source) + prompt = (choice.model_extra or {}).get("prompt_token_ids") + if prompt is None: + prompt = (exchange.response.model_extra or {}).get("prompt_token_ids") + output = (choice.model_extra or {}).get("token_ids") + if isinstance(prompt, list) and isinstance(output, list): + return prompt, output + prompt, output, _ = _source_native_record(source) + return prompt, output + + +def _partial_native_context(history: History | LegacyHistory) -> list[object]: + from ._history import _retains_output_suffix + + if isinstance(history, (ChatCompletionsHistory, AnthropicMessagesHistory)): + sources: Sequence[object] = history.message_sources + elif isinstance(history, ResponsesHistory): + sources = history.input_sources + else: + return [] + sampled = [ + source + for source in sources + if source is not None and _source_is_sampled(source) + ] + if len(sampled) < 2: + return [] + final_prompt, _ = _source_native_prefix(sampled[-1]) + if final_prompt is None: + return [] + partial = [] + for source in sampled[:-1]: + prompt, output = _source_native_prefix(source) + if prompt is None or output is None: + continue + if ( + final_prompt[: len(prompt)] == prompt + and final_prompt[len(prompt) : len(prompt) + len(output)] == output + ): + continue + if _retains_output_suffix(prompt, output, final_prompt): + partial.append(source) + return partial + + +def _certify_copied_context( + tokenized: TokenizedHistory, + trace: _HistoryTokenizationTrace, + copied: Sequence[object], + prior: Sequence[tuple[TokenizedHistory, _HistoryTokenizationTrace]], +) -> None: + """A rendered copy may keep output provenance, never its old prediction LP.""" + copied_keys = {_sampled_source_key(source) for source in copied} + positions: dict[_SampledSourceKey, list[int]] = {} + for index, key in enumerate(trace.source_keys): + if key is not None: + positions.setdefault(key, []).append(index) + for key, offsets in positions.items(): + source = trace.sources[key] + prompt, output, logprobs = _source_native_record(source) + if prompt is None or output is None: + continue + end = len(prompt) + len(output) + complete = offsets == list(range(len(prompt), end)) and tokenized.tokens[ + :end + ] == [*prompt, *output] + if complete: + continue + if key not in copied_keys or not _complete_source_is_represented( + source, prompt, output, logprobs, prior + ): + raise ValueError( + "Recorded sampled tokens do not retain their original native conditioning" + ) + # The unchanged original source remains trainable in an earlier result; + # these tokens are a copy under a different prompt, not another draw. + for index in offsets: + tokenized.flags[index] &= ~TokenFlag.SAMPLED + tokenized.logprobs[index] = math.nan + trace.source_keys[index] = None + trace.validate(tokenized) + + def _tokenize_exact_projected_chat_history( history: ChatCompletionsHistory, *, @@ -4126,6 +4324,8 @@ def _tokenize_exact_projected_chat_history( | None = None, projection_validated: bool = False, _trace: _TraceBuilder | None = None, + _strict_sources: bool = False, + _prior: Sequence[tuple[TokenizedHistory, _HistoryTokenizationTrace]] = (), ) -> TokenizedHistory | None: if not projection_validated and not _history_matches_projection(history): return None @@ -4145,9 +4345,19 @@ def _tokenize_exact_projected_chat_history( if not sampled_sources: return None + # A later-prompt lookup must not repeatedly decode that source's output LPs. + # Keep validated records only for this assembly; never across render calls. + records: dict[int, tuple[list[int] | None, list[int] | None, list[float]]] = {} + + def record( + source: object, + ) -> tuple[list[int] | None, list[int] | None, list[float]]: + if id(source) not in records: + records[id(source)] = _chat_source_record(source) + return records[id(source)] + final_source = sampled_sources[-1] - final_prompt = _chat_source_prompt_tokens(final_source) - final_output, final_logprobs = _chat_source_full_tokens(final_source) + final_prompt, final_output, final_logprobs = record(final_source) if final_prompt is None or final_output is None: return None final_key = _sampled_source_key(final_source) @@ -4223,8 +4433,7 @@ def _tokenize_exact_projected_chat_history( ] sources: dict[_SampledSourceKey, object] = {final_key: final_source} for index, source in enumerate(sampled_sources[:-1]): - prompt = _chat_source_prompt_tokens(source) - output, output_logprobs = _chat_source_full_tokens(source) + prompt, output, output_logprobs = record(source) if ( prompt is None or output is None @@ -4235,8 +4444,7 @@ def _tokenize_exact_projected_chat_history( ( evidence for later_source in sampled_sources[index + 1 :] - if (later_prompt := _chat_source_prompt_tokens(later_source)) - is not None + if (later_prompt := record(later_source)[0]) is not None and ( evidence := _retained_output_suffix( prompt=prompt, @@ -4261,9 +4469,55 @@ def _tokenize_exact_projected_chat_history( ] * len(retained_ids) logprobs[start:end] = retained_logprobs source_key = _sampled_source_key(source) - if _source_stop_evidence(source, source_key)[0] == "length": + if _strict_sources and retained_ids != output: + if not _complete_source_is_represented( + source, prompt, output, output_logprobs, _prior + ): + raise ValueError( + "A copied response suffix has different native conditioning; " + "its complete original sampled occurrence must be represented " + "in the selected trajectory before it can be used as context" + ) + flags[start:end] = [ + TokenFlag.EXACT | TokenFlag.ASSISTANT | TokenFlag.OUTPUT + ] * len(retained_ids) + logprobs[start:end] = [math.nan] * len(retained_ids) + records.clear() # A custom STOP decoder may change source objects. + stop_count = _sampled_stop_suffix( + output, source=source, source_key=source_key, tokenizer=tokenizer + ) + for offset in range(max(start, end - stop_count), end): + flags[offset] |= TokenFlag.STOP + boundary = (length_stop_boundaries or {}).get(source_key) + if boundary is not None: + tail_end = end + len(boundary.tail) + next_prompt = record(sampled_sources[index + 1])[0] + if not boundary.tail or next_prompt != [ + *final_prompt[:end], + *boundary.tail, + *boundary.following, + ]: + return None + stop_kind = _source_stop_evidence(source, source_key)[0] + boundary_flags = TokenFlag.EXACT | TokenFlag.ASSISTANT + if stop_kind == "stop": + boundary_flags |= TokenFlag.OUTPUT + flags[end:tail_end] = [boundary_flags] * len(boundary.tail) + flags[tail_end - 1] = ( + TokenFlag.EXACT + | TokenFlag.STOP + | ( + TokenFlag.ASSISTANT | TokenFlag.OUTPUT + if stop_kind == "stop" + else TokenFlag(0) + ) + ) + sources[source_key] = source + continue + stop_kind = _source_stop_evidence(source, source_key)[0] + if stop_kind == "length" or source_key in (length_stop_boundaries or {}): boundary = (length_stop_boundaries or {}).get(source_key) - next_prompt = _chat_source_prompt_tokens(sampled_sources[index + 1]) + next_prompt = record(sampled_sources[index + 1])[0] if boundary is not None and next_prompt is not None: rendered_boundary = [*boundary.tail, *boundary.following] native_boundary = next_prompt[end:] @@ -4273,14 +4527,15 @@ def _tokenize_exact_projected_chat_history( extra > 0 and native_boundary[extra:] == rendered_boundary and callable(decode) - and decode(native_boundary[:extra]).isspace() ): - # Services may insert whitespace before a truncated turn's - # proven stop tail. Keep those served, nonsampled tokens. - boundary = _RenderedLengthStopBoundary( - tail=(*native_boundary[:extra], *boundary.tail), - following=boundary.following, - ) + records.clear() # Never reuse records across user callbacks. + if decode(native_boundary[:extra]).isspace(): + # Services may insert whitespace before a truncated turn's + # proven stop tail. Keep those served, nonsampled tokens. + boundary = _RenderedLengthStopBoundary( + tail=(*native_boundary[:extra], *boundary.tail), + following=boundary.following, + ) boundary_end = ( end + len(boundary.tail) + len(boundary.following) if boundary is not None @@ -4299,10 +4554,19 @@ def _tokenize_exact_projected_chat_history( # output and renderer-proven boundary, render the stop. return None tail_end = end + len(boundary.tail) - flags[end:tail_end] = [TokenFlag.EXACT | TokenFlag.ASSISTANT] * len( - boundary.tail + boundary_flags = TokenFlag.EXACT | TokenFlag.ASSISTANT + if stop_kind == "stop": + boundary_flags |= TokenFlag.OUTPUT + flags[end:tail_end] = [boundary_flags] * len(boundary.tail) + flags[tail_end - 1] = ( + TokenFlag.EXACT + | TokenFlag.STOP + | ( + TokenFlag.ASSISTANT | TokenFlag.OUTPUT + if stop_kind == "stop" + else TokenFlag(0) + ) ) - flags[tail_end - 1] = TokenFlag.EXACT | TokenFlag.STOP source_keys[start:end] = [source_key] * len(retained_ids) sources[source_key] = source if history.model is None: @@ -4326,6 +4590,134 @@ def _tokenize_exact_projected_chat_history( return tokenized +def _tokenize_recorded_chat_boundaries( + history: ChatCompletionsHistory, + messages: list[dict[str, Any]], + *, + tokenizer: Tokenizer, + render: _ChatRender, + _trace: _TraceBuilder | None, + _prior: Sequence[tuple[TokenizedHistory, _HistoryTokenizationTrace]] = (), +) -> TokenizedHistory | None: + """Reuse complete native spans; encode only unrecorded turn boundaries. + + This does not repartition histories or infer flags for request-owned assistant + messages. Unsupported render/decode capabilities retain the ordinary path. + """ + decode = getattr(tokenizer, "decode", None) + if not callable(decode) or not messages or messages[-1].get("role") != "assistant": + return None + entries: list[tuple[int, object, list[int], list[int], list[float]]] = [] + seen: set[_SampledSourceKey] = set() + for index, (message, source) in enumerate( + zip(messages, history.message_sources, strict=True) + ): + if message.get("role") != "assistant": + continue + if source is None or not _source_is_sampled(source): + return None + key = _sampled_source_key(source) + prompt, output, logprobs = _chat_source_record(source) + if key in seen or prompt is None or output is None: + return None + seen.add(key) + entries.append((index, source, prompt, output, logprobs)) + if not entries: + return None + final_prompt = entries[-1][2] + for ordinal, (index, source, prompt, output, logprobs) in enumerate(entries[:-1]): + retained = _retained_output_suffix( + prompt=prompt, output=output, logprobs=logprobs, later_prompt=final_prompt + ) + if retained is not None and retained[0] != output: + if not _complete_source_is_represented( + source, prompt, output, logprobs, _prior + ): + raise ValueError( + "A copied response suffix requires its complete original sampled occurrence in the selected trajectory" + ) + entries[ordinal] = (index, source, prompt, retained[0], retained[1]) + # The same canonical history must contain every original conditioning edge. + # Text equivalence is insufficient: these comparisons are exact native IDs. + for (_, _, prompt, output, _), (_, _, next_prompt, _, _) in zip( + entries, entries[1:] + ): + if ( + next_prompt[: len(prompt)] != prompt + or next_prompt[len(prompt) : len(prompt) + len(output)] != output + ): + return None + boundaries: dict[_SampledSourceKey, _RenderedLengthStopBoundary] = {} + terminators = _terminator_ids(tokenizer) + if not terminators: + return None + for ordinal, (index, source, prompt, output, _) in enumerate(entries): + key = _sampled_source_key(source) + stop, _ = _source_stop_evidence(source, key) + if stop not in {"stop", "length"}: + return None + if stop == "stop" and _sampled_stop_suffix( + output, source=source, source_key=key, tokenizer=tokenizer + ): + continue + try: + body = decode( + output, + skip_special_tokens=False, + clean_up_tokenization_spaces=False, + ) + generation = render(messages[:index], add_generation_prompt=True) + completed = render(messages[: index + 1], add_generation_prompt=False) + # Only the actually sampled body anchors the tail. Literal content, + # tool JSON and reasoning are never searched for or re-tokenized. + if not isinstance(body, str) or not completed.startswith(generation + body): + return None + suffix = completed[len(generation) + len(body) :] + tail = _ids(tokenizer(suffix, add_special_tokens=False)) if suffix else [] + stops = [i for i, token in enumerate(tail) if token in terminators] + if len(stops) != 1: + return None + terminator = stops[0] + trailing = decode( + tail[terminator + 1 :], + skip_special_tokens=False, + clean_up_tokenization_spaces=False, + ) + if not isinstance(trailing, str) or trailing and not trailing.isspace(): + return None + following: list[int] = [] + if ordinal + 1 < len(entries): + next_index, _, next_prompt, _, _ = entries[ordinal + 1] + next_generation = render( + messages[:next_index], add_generation_prompt=True + ) + if not next_generation.startswith(completed): + return None + gap = suffix + next_generation[len(completed) :] + gap_ids = _ids(tokenizer(gap, add_special_tokens=False)) + if ( + gap_ids[: len(tail)] != tail + or next_prompt[len(prompt) + len(output) :] != gap_ids + ): + return None + following = gap_ids[len(tail) :] + boundaries[key] = _RenderedLengthStopBoundary( + tail=tuple(tail[: terminator + 1]), + following=tuple([*tail[terminator + 1 :], *following]), + ) + except (TypeError, KeyError, NotImplementedError): + return None + return _tokenize_exact_projected_chat_history( + history, + tokenizer=tokenizer, + length_stop_boundaries=boundaries, + projection_validated=True, + _trace=_trace, + _strict_sources=True, + _prior=_prior, + ) + + def _chat_message_parts(message: Mapping[str, object]) -> list[tuple[str, str]]: parts: list[tuple[str, str]] = [] reasoning = message.get("reasoning") @@ -4532,58 +4924,6 @@ def _source_covers_complete_sampled_message( ) == normalize_chat_message(projected[0]) -def _preserve_literal_thinking_off_content( - history: ChatCompletionsHistory, - messages: list[dict[str, Any]], - template: object, - kwargs: Mapping[str, object], -) -> None: - # This Qwen3.5 template treats any in unstructured content as a - # reasoning separator, even with thinking disabled. Restrict the render-copy - # adaptation to its exact preserved template; other templates may interpret - # an empty reasoning_content field differently. - if ( - not isinstance(template, str) - or sha256(template.encode()).hexdigest() - != "098047d425a6673b1fe1a82a197a481616e53a283beaa8cb76cbb74d38ca6644" - or kwargs.get("enable_thinking") is not False - or kwargs.get("preserve_thinking") is not True - ): - return - for message, source in zip(messages, history.message_sources, strict=True): - if ( - source is None - or not isinstance(source.exchange, ChatCompletionsExchange) - or source.choice_index is None - or message.get("role") != "assistant" - or not isinstance(content := message.get("content"), str) - or "" not in content - ): - continue - request_kwargs = source.exchange.request.get("chat_template_kwargs") - if ( - not isinstance(request_kwargs, Mapping) - or request_kwargs.get("enable_thinking") is not False - ): - continue - choice = _chat_choice(source) - # Visible-only histories may omit structured reasoning present in the - # source response. Preserve both that source and normalized aliases. - if any( - value is not None and not (isinstance(value, str) and value == "") - for value in ( - message.get("reasoning"), - message.get("reasoning_content"), - _field(choice.message, "reasoning"), - _field(choice.message, "reasoning_content"), - ) - ): - continue - prompt, output, _ = _chat_choice_tokens(choice, source.exchange.response) - if prompt is not None and output is not None: - message["reasoning_content"] = "" - - def _tokenize_chat_view( history: ChatCompletionsHistory, *, @@ -4592,7 +4932,9 @@ def _tokenize_chat_view( chat_template: str | None, chat_template_kwargs: Mapping[str, object] | None, _projection_matches: bool | None = None, + _recorded_boundaries: bool = False, _trace: _TraceBuilder | None = None, + _prior: Sequence[tuple[TokenizedHistory, _HistoryTokenizationTrace]] = (), ) -> TokenizedHistory: _validate_history_sources(history) config = ( @@ -4630,7 +4972,6 @@ def _tokenize_chat_view( **default_chat_template_kwargs_for_template(template), **explicit_kwargs, } - _preserve_literal_thinking_off_content(history, messages, template, kwargs) ends_with_assistant = bool(messages) and messages[-1].get("role") == "assistant" segmented = False @@ -4676,6 +5017,17 @@ def render_text( add_generation_prompt=add_generation_prompt, ) + if _recorded_boundaries: + if recorded := _tokenize_recorded_chat_boundaries( + history, + messages, + tokenizer=resolved_tokenizer, + render=render_text, + _trace=_trace, + _prior=_prior, + ): + return recorded + prefix_render_cache = _PrefixChatRenderCache(render_normalized_text) def segmented_render( @@ -6655,6 +7007,7 @@ def _tokenize_history( chat_template: str | None, chat_template_kwargs: Mapping[str, object] | None, _trace: _TraceBuilder | None = None, + _prior: Sequence[tuple[TokenizedHistory, _HistoryTokenizationTrace]] = (), _projection_validated: bool = False, ) -> TokenizedHistory: if isinstance(history, LegacyHistory): @@ -6703,7 +7056,11 @@ def _tokenize_history( ) if isinstance(history, ResponsesHistory) and not needs_render: if exact := _tokenize_exact_responses_history( - history, base_model=base_model, tokenizer=tokenizer, _trace=_trace + history, + base_model=base_model, + tokenizer=tokenizer, + _trace=_trace, + _prior=_prior, ): return exact if isinstance(history, ChatCompletionsHistory): @@ -6721,6 +7078,8 @@ def _tokenize_history( _projection_validated or render_state.projection_matches is True ), _trace=_trace, + _strict_sources=True, + _prior=_prior, ) ) ): @@ -6734,6 +7093,13 @@ def _tokenize_history( _projection_matches=( True if _projection_validated else render_state.projection_matches ), + _prior=_prior, + _recorded_boundaries=( + (has_length_stop or needs_synthetic_stop) + and not override_requires_render + and not render_state.context_changed + and (_projection_validated or render_state.projection_matches is True) + ), _trace=_trace, ) if isinstance(history, AnthropicMessagesHistory) and needs_render: @@ -6805,8 +7171,51 @@ def tokenize_history( chat_template: str | None, chat_template_kwargs: Mapping[str, object] | None, _trace: _TraceBuilder | None = None, + _prior: Sequence[tuple[TokenizedHistory, _HistoryTokenizationTrace]] = (), _projection_validated: bool = False, + _context_sources: Sequence[object] | None = None, ) -> TokenizedHistory: + copied = ( + list(_context_sources) + if _context_sources is not None + else _partial_native_context(history) + ) + if copied: + history = cast(History, history) + _validate_history_sources(history) + state = None if _projection_validated else _history_render_state(history) + unchanged = _projection_validated or ( + state is not None + and not state.context_changed + and ( + state.projection_matches is True + or state.projection_matches is None + and _history_matches_projection(history) + ) + ) + override = ( + chat_template is not None + and chat_template != getattr(history, "chat_template", None) + ) or ( + chat_template_kwargs is not None + and dict(chat_template_kwargs) + != (getattr(history, "chat_template_kwargs", None) or {}) + ) + if not unchanged or override: + copied = [] + for source in copied: + prompt, output, logprobs = _source_native_record(source) + if ( + prompt is None + or output is None + or not _complete_source_is_represented( + source, prompt, output, logprobs, _prior + ) + ): + raise ValueError( + "A copied response suffix requires its complete original sampled occurrence in the selected trajectory" + ) + trace_builder = _trace or (_TraceBuilder() if copied else None) tokenized = _tokenize_history( history, model=model, @@ -6814,9 +7223,16 @@ def tokenize_history( tokenizer=tokenizer, chat_template=chat_template, chat_template_kwargs=chat_template_kwargs, - _trace=_trace, + _trace=trace_builder, + _prior=_prior, _projection_validated=_projection_validated, ) + if copied: + if trace_builder is None or trace_builder.trace is None: + raise ValueError( + "Copied native context requires a complete tokenization source trace" + ) + _certify_copied_context(tokenized, trace_builder.trace, copied, _prior) # Internal protocol conversion is an implementation detail. The source is # always the public history view the caller asked to tokenize. if not isinstance( @@ -6877,8 +7293,13 @@ def tokenize_trajectory( raise ValueError( f"Trajectory tokenization requires exactly one history; found {len(histories)}" ) - tokenized = [ - tokenize_history( + context_sources = [_partial_native_context(history) for history in histories] + track_context = len(histories) > 1 and any(context_sources) + prior: list[tuple[TokenizedHistory, _HistoryTokenizationTrace]] = [] + tokenized = [] + for history, copied in zip(histories, context_sources, strict=True): + trace = _TraceBuilder() if track_context else None + result = tokenize_history( history, model=model if isinstance(history, LegacyHistory) else history.model, base_model=base_model, @@ -6886,9 +7307,13 @@ def tokenize_trajectory( chat_template=chat_template, chat_template_kwargs=chat_template_kwargs, _projection_validated=not isinstance(history, LegacyHistory), + _trace=trace, + _prior=prior, + _context_sources=copied, ) - for history in histories - ] + tokenized.append(result) + if trace is not None and trace.trace is not None: + prior.append((result, trace.trace)) if not multi_history: return _materialize_trajectory(tokenized[0], trajectory) return TokenizedMultiHistoryTrajectory( @@ -6929,6 +7354,7 @@ def _tokenize_trajectory_with_trace( chat_template_kwargs=chat_template_kwargs, _trace=trace_builder, _projection_validated=True, + _prior=list(zip(tokenized_histories, traces, strict=True)), ) if trace_builder.trace is None: raise AssertionError("Exchange tokenization did not produce a source trace") diff --git a/tests/unit/test_literal_reasoning_content.py b/tests/unit/test_literal_reasoning_content.py index 0df6826a9..8561dc5fa 100644 --- a/tests/unit/test_literal_reasoning_content.py +++ b/tests/unit/test_literal_reasoning_content.py @@ -215,7 +215,9 @@ def test_unconfigured_template_receives_the_same_correction(): def test_inline_operation_as_raw_or_comment_text_is_not_rewritten(wrapper): from art_inference.chat_template import _QWEN_INLINE_REASONING - operation = _QWEN_INLINE_REASONING.search(_TEMPLATE).group() + match = _QWEN_INLINE_REASONING.search(_TEMPLATE) + assert match is not None + operation = match.group() template = wrapper[0] + operation + wrapper[1] assert chat_template_with_preserved_thinking(template) == template @@ -224,6 +226,7 @@ def test_other_structured_reasoning_condition_is_not_rewritten(): gate = "{% if preserve_thinking and message.role == 'assistant' %}{{ message.reasoning_content }}{% endif %}" template = _TEMPLATE + "{% for message in messages %}" + gate + "{% endfor %}" fixed = chat_template_with_preserved_thinking(template) + assert isinstance(fixed, str) assert gate in fixed messages = [ _USER, @@ -236,7 +239,9 @@ def test_other_structured_reasoning_condition_is_not_rewritten(): def test_inline_operation_inside_quoted_expression_is_literal(): from art_inference.chat_template import _QWEN_INLINE_REASONING - operation = _QWEN_INLINE_REASONING.search(_TEMPLATE).group().replace("\n", " ") + match = _QWEN_INLINE_REASONING.search(_TEMPLATE) + assert match is not None + operation = match.group().replace("\n", " ") template = '{{ "' + operation + '" }}' fixed = chat_template_with_preserved_thinking(template) assert fixed == template @@ -249,7 +254,9 @@ def test_inline_operation_inside_quoted_expression_is_literal(): def test_mixed_executable_and_literal_operations_only_changes_executable(wrapper): from art_inference.chat_template import _QWEN_INLINE_REASONING - operation = _QWEN_INLINE_REASONING.search(_TEMPLATE).group().replace("\n", " ") + match = _QWEN_INLINE_REASONING.search(_TEMPLATE) + assert match is not None + operation = match.group().replace("\n", " ") literal = wrapper[0] + operation + wrapper[1] template = _TEMPLATE + literal fixed = chat_template_with_preserved_thinking(template) diff --git a/tests/unit/trajectories/test_literal_thinking_off.py b/tests/unit/trajectories/test_literal_thinking_off.py index 03e567e21..ca423cd41 100644 --- a/tests/unit/trajectories/test_literal_thinking_off.py +++ b/tests/unit/trajectories/test_literal_thinking_off.py @@ -145,9 +145,6 @@ def test_native_thinking_off_retains_literal_content( # The pre-fix history path misrenders literal content even when later native # token splicing can recover the terminal output. with monkeypatch.context() as patch: - patch.setattr( - _tokenize, "_preserve_literal_thinking_off_content", lambda *args: None - ) patch.setattr( _tokenize, "chat_template_with_preserved_thinking", lambda value: value ) @@ -186,9 +183,7 @@ def test_native_thinking_off_retains_literal_content( "visible_only", ], ) -def test_unrelated_histories_keep_original_rendering( - case: str, monkeypatch: pytest.MonkeyPatch -) -> None: +def test_literal_content_is_not_inferred_from_source_thinking_mode(case: str) -> None: history, tokenizer = _history( thinking=True if case == "source_on" @@ -225,17 +220,19 @@ def test_unrelated_histories_keep_original_rendering( if case == "visible_only": cast(dict[str, Any], history.messages[-1]).pop("reasoning") original = history.model_dump(mode="python") - candidate = _outcome(history, tokenizer) - calls = deepcopy(tokenizer.calls) - tokenizer.calls.clear() - with monkeypatch.context() as patch: - patch.setattr( - _tokenize, "_preserve_literal_thinking_off_content", lambda *args: None + _outcome(history, tokenizer) + # Plain content stays literal independently of recorded/current thinking mode + # and whether the message has complete native token metadata. Structured + # reasoning remains a separate field on the render copy. + assert tokenizer.calls[0][-1]["content"] == _LITERAL + assert _LITERAL in tokenizer.rendered[0] + if case in {"structured", "alias"}: + assert ( + tokenizer.calls[0][-1].get( + "reasoning_content", tokenizer.calls[0][-1].get("reasoning") + ) + == "explicit reasoning" ) - baseline = _outcome(history, tokenizer) - assert candidate == baseline - assert len(calls) == len(tokenizer.calls) - assert calls[0] == tokenizer.calls[0] assert history.model_dump(mode="python") == original @@ -406,9 +403,6 @@ def observe(*args: Any, **kwargs: Any) -> tr.TokenizedHistory | None: monkeypatch.setattr(_tokenize, "_tokenize_exact_projected_chat_history", observe) with monkeypatch.context() as patch: - patch.setattr( - _tokenize, "_preserve_literal_thinking_off_content", lambda *args: None - ) patch.setattr( _tokenize, "chat_template_with_preserved_thinking", lambda value: value ) diff --git a/tests/unit/trajectories/test_recorded_boundaries.py b/tests/unit/trajectories/test_recorded_boundaries.py new file mode 100644 index 000000000..ef5db030f --- /dev/null +++ b/tests/unit/trajectories/test_recorded_boundaries.py @@ -0,0 +1,555 @@ +from __future__ import annotations + +import math +from typing import Any, cast + +import pytest +from test_tokenize import _character_template_history + +from art.trajectories import TokenFlag, first_occurrence_masks +from art.trajectories import _tokenize as module + + +@pytest.fixture(autouse=True) +def restore_warning_state(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(module, "_WARNED_PREFIX_RETOKENIZATION", False) + + +def assert_same(left: Any, right: Any) -> None: + assert left.history is right.history + assert left.model == right.model + assert left.tokens == right.tokens + assert left.flags == right.flags + assert len(left.logprobs) == len(right.logprobs) + assert all( + a == b or math.isnan(a) and math.isnan(b) + for a, b in zip(left.logprobs, right.logprobs, strict=True) + ) + for flag in ( + TokenFlag.SAMPLED, + TokenFlag.OUTPUT, + TokenFlag.ASSISTANT, + TokenFlag.STOP, + ): + assert first_occurrence_masks([left], where=flag) == first_occurrence_masks( + [right], where=flag + ) + + +@pytest.mark.parametrize("terminal_sampled_stop", [False, True]) +def test_recorded_length_boundaries_do_not_reencode_sampled_content( + monkeypatch: pytest.MonkeyPatch, terminal_sampled_stop: bool +) -> None: + history, tokenizer, _ = _character_template_history( + terminal_sampled_stop=terminal_sampled_stop + ) + original = history.model_dump(mode="python") + helper = module._tokenize_recorded_chat_boundaries + admissions = [] + + def observe(*args: Any, **kwargs: Any) -> Any: + value = helper(*args, **kwargs) + admissions.append(value) + return value + + monkeypatch.setattr(module, "_tokenize_recorded_chat_boundaries", observe) + rendered = [] + original_render = tokenizer.apply_chat_template + + def render(*args: Any, **kwargs: Any) -> Any: + rendered.append(kwargs.get("tokenize", True)) + return original_render(*args, **kwargs) + + monkeypatch.setattr(tokenizer, "apply_chat_template", render) + tokenized = history.tokenize(tokenizer=tokenizer) + assert admissions == [tokenized] + assert rendered and not any(rendered) + monkeypatch.setattr( + module, "_tokenize_recorded_chat_boundaries", lambda *args, **kwargs: None + ) + baseline = history.tokenize(tokenizer=tokenizer) + assert_same(tokenized, baseline) + assert history.model_dump(mode="python") == original + + +@pytest.mark.parametrize("change", ["missing_tail", "reasoning", "override", "edited"]) +def test_unproved_boundaries_preserve_existing_path( + monkeypatch: pytest.MonkeyPatch, change: str +) -> None: + history, tokenizer, _ = _character_template_history( + omit_length_tail=change == "missing_tail", + length_reasoning="not in native output" if change == "reasoning" else None, + ) + kwargs = ( + {"chat_template": "explicit caller template"} if change == "override" else {} + ) + if change == "edited": + history.messages[3]["content"] = "edited" + helper = module._tokenize_recorded_chat_boundaries + admissions = [] + + def observe(*args: Any, **kwargs: Any) -> Any: + value = helper(*args, **kwargs) + admissions.append(value) + return value + + def result() -> Any: + try: + return history.tokenize(tokenizer=tokenizer, **kwargs) + except (ValueError, AssertionError) as error: + return type(error), str(error) + + monkeypatch.setattr(module, "_tokenize_recorded_chat_boundaries", observe) + candidate = result() + if change != "reasoning": + assert not any(admissions) + monkeypatch.setattr( + module, "_tokenize_recorded_chat_boundaries", lambda *args, **kwargs: None + ) + baseline = result() + if isinstance(candidate, tuple): + assert candidate == baseline + else: + assert_same(candidate, baseline) + + +@pytest.mark.parametrize("tool_position", [0, 1]) +def test_recorded_tool_boundaries_preserve_native_conditioning( + monkeypatch: pytest.MonkeyPatch, tool_position: int +) -> None: + from copy import deepcopy + import json + + from openai.types.chat import ChatCompletion, ChatCompletionMessageParam + from test_tokenize import _CharacterTemplateTokenizer, _chat_exchange + + import art.trajectories as tr + + class Tokenizer(_CharacterTemplateTokenizer): + def apply_chat_template( + self, + messages: list[dict[str, Any]], + *, + tokenize: bool = True, + add_generation_prompt: bool, + **kwargs: Any, + ) -> str | list[int]: + del kwargs + text = "" + for message in messages: + text += message["role"] + ":" + str(message.get("content") or "") + if message.get("tool_calls"): + text += json.dumps(message["tool_calls"], sort_keys=True) + if message["role"] == "assistant": + text += "§" + if add_generation_prompt: + text += "assistant:" + return self._encode(text) if tokenize else text + + tokenizer = Tokenizer() + exchanges = [] + messages: list[dict[str, Any]] = [] + expected_spans = [] + for index in range(2): + messages.append({"role": "user", "content": f"query{index}"}) + prompt = tokenizer.apply_chat_template(messages, add_generation_prompt=True) + message = {"role": "assistant", "content": "answer"} + if index == tool_position: + message = { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "public_call", + "type": "function", + "function": {"name": "lookup", "arguments": '{"x":1}'}, + } + ], + } + completion = tokenizer.apply_chat_template( + [*messages, message], add_generation_prompt=False + ) + assert isinstance(prompt, list) and isinstance(completion, list) + output = completion[len(prompt) :] + if index == tool_position: + output = output[:-1] # Server stopped on a tool call without emitting EOS. + exchange = _chat_exchange(prompt, output, offset=index) + exchange.request["messages"] = cast( + list[ChatCompletionMessageParam], deepcopy(messages) + ) + payload = exchange.response.model_dump(mode="python") + payload["choices"][0]["message"] = message + payload["choices"][0]["finish_reason"] = ( + "tool_calls" if index == tool_position else "stop" + ) + exchange.response = ChatCompletion.model_validate(payload) + exchanges.append(exchange) + messages.append(message) + expected_spans.append((prompt, output)) + trajectory = tr.Trajectory( + exchanges=tr.TrajectoryExchanges(chat_completions=exchanges) + ) + before = trajectory.model_dump(mode="python") + result = trajectory.tokenize(tokenizer=tokenizer, multi_history=True) + assert len(result.histories) == 1 + tokenized = result.histories[0] + for prompt, output in expected_spans: + assert tokenized.tokens[: len(prompt)] == prompt + assert tokenized.tokens[len(prompt) : len(prompt) + len(output)] == output + assert all( + flag & TokenFlag.SAMPLED + for flag in tokenized.flags[len(prompt) : len(prompt) + len(output)] + ) + assert sum(bool(flag & TokenFlag.STOP) for flag in tokenized.flags) == 2 + assert trajectory.model_dump(mode="python") == before + monkeypatch.setattr( + module, "_tokenize_recorded_chat_boundaries", lambda *args, **kwargs: None + ) + try: + baseline = trajectory.tokenize(tokenizer=tokenizer, multi_history=True) + except ValueError as error: + assert "boundary" in str(error) or "prefix" in str(error) + else: + old = baseline.histories[0] + if tool_position == 0: + # The former text replacement deleted the served, nonsampled EOS: + # its returned second response no longer had its recorded prompt. + prompt, _ = expected_spans[1] + assert old.tokens[: len(prompt)] != prompt + else: + # The template owns this EOS; it must not become a sampled token. + assert old.tokens == tokenized.tokens[:-1] + tool_prompt, tool_output = expected_spans[tool_position] + stop_position = len(tool_prompt) + len(tool_output) + assert tokenized.flags[stop_position] & TokenFlag.STOP + assert not tokenized.flags[stop_position] & TokenFlag.SAMPLED + + +@pytest.mark.parametrize("footer", ["footer§", "user-owned footer"]) +def test_custom_footer_is_not_inferred_as_an_assistant_boundary( + monkeypatch: pytest.MonkeyPatch, footer: str +) -> None: + history, tokenizer, _ = _character_template_history() + original = tokenizer.apply_chat_template + + def render( + messages: Any, + *, + tokenize: bool = True, + add_generation_prompt: bool, + **kwargs: Any, + ) -> Any: + text = original( + messages, + tokenize=False, + add_generation_prompt=add_generation_prompt, + **kwargs, + ) + if not add_generation_prompt and messages[-1]["role"] == "assistant": + assert isinstance(text, str) + text += footer + assert isinstance(text, str) + return tokenizer._encode(text) if tokenize else text + + monkeypatch.setattr(tokenizer, "apply_chat_template", render) + helper = module._tokenize_recorded_chat_boundaries + outcomes = [] + + def observe(*args: Any, **kwargs: Any) -> Any: + value = helper(*args, **kwargs) + outcomes.append(value) + return value + + monkeypatch.setattr(module, "_tokenize_recorded_chat_boundaries", observe) + try: + actual = history.tokenize(tokenizer=tokenizer) + except ValueError as error: + actual = type(error), str(error) + assert outcomes == [None] + monkeypatch.setattr( + module, "_tokenize_recorded_chat_boundaries", lambda *args, **kwargs: None + ) + try: + expected = history.tokenize(tokenizer=tokenizer) + except ValueError as error: + expected = type(error), str(error) + if isinstance(actual, tuple): + assert actual == expected + else: + assert_same(actual, expected) + + +@pytest.mark.parametrize("logprob", [-0.3, math.nan, 1e100]) +def test_copied_suffix_is_context_not_a_new_sampled_edge(logprob: float) -> None: + from test_tokenize import _chat_exchange + + import art.trajectories as tr + + first = _chat_exchange([1], [2, 3]) + recorded = first.response.choices[0].logprobs + assert recorded is not None and recorded.content is not None + recorded.content[-1].logprob = logprob + second = _chat_exchange([1, 3, 4], [5], offset=1) + trajectory = tr.Trajectory( + exchanges=tr.TrajectoryExchanges(chat_completions=[first, second]) + ) + original = trajectory.model_dump_json() + tokenized = trajectory.tokenize(multi_history=True) + assert len(tokenized.histories) == 2 + original_result, copied_result = tokenized.histories + assert original_result.tokens == [1, 2, 3] + assert original_result.flags == [ + TokenFlag.EXACT, + (TokenFlag.EXACT | TokenFlag.SAMPLED | TokenFlag.ASSISTANT | TokenFlag.OUTPUT), + (TokenFlag.EXACT | TokenFlag.SAMPLED | TokenFlag.ASSISTANT | TokenFlag.OUTPUT), + ] + assert ( + original_result.logprobs[2] == logprob + or math.isnan(original_result.logprobs[2]) + and math.isnan(logprob) + ) + assert copied_result.tokens == [1, 3, 4, 5] + assert copied_result.flags == [ + TokenFlag.EXACT, + TokenFlag.EXACT | TokenFlag.ASSISTANT | TokenFlag.OUTPUT, + TokenFlag.EXACT, + (TokenFlag.EXACT | TokenFlag.SAMPLED | TokenFlag.ASSISTANT | TokenFlag.OUTPUT), + ] + assert math.isnan(copied_result.logprobs[1]) + assert first_occurrence_masks(tokenized.histories, where=TokenFlag.SAMPLED) == [ + [False, True, True], + [False, False, False, True], + ] + assert trajectory.model_dump_json() == original + standalone = trajectory.histories()[1] + assert isinstance(standalone, tr.ChatCompletionsHistory) + with pytest.raises(ValueError, match="complete original sampled occurrence"): + standalone.tokenize() + + +def test_length_copy_keeps_proven_synthetic_boundary_flags() -> None: + from openai.types.chat import ChatCompletion + from test_tokenize import _CharacterTemplateTokenizer, _chat_exchange + + import art.trajectories as tr + + class Tokenizer(_CharacterTemplateTokenizer): + def apply_chat_template( + self, + messages: Any, + *, + tokenize: bool = True, + add_generation_prompt: bool, + **kwargs: Any, + ) -> Any: + text = "".join( + str(message.get("reasoning") or message.get("reasoning_content") or "") + + str(message.get("content") or "") + + ("§" if message["role"] == "assistant" else "") + for message in messages + ) + return self._encode(text) if tokenize else text + + tokenizer = Tokenizer() + prompt = tokenizer._encode("turn 0") + first = _chat_exchange(prompt, tokenizer._encode("ranswer")) + payload = first.response.model_dump(mode="python") + payload["choices"][0]["message"]["reasoning_content"] = "r" + payload["choices"][0]["finish_reason"] = "length" + first.response = ChatCompletion.model_validate(payload) + next_prompt = tokenizer._encode("turn 0answer§turn 1") + second = _chat_exchange(next_prompt, tokenizer._encode("answer§"), offset=1) + trajectory = tr.Trajectory( + exchanges=tr.TrajectoryExchanges(chat_completions=[first, second]) + ) + result = trajectory.tokenize(tokenizer=tokenizer, multi_history=True) + assert len(result.histories) == 2 + original, copied = result.histories + assert original.tokens == tokenizer._encode("turn 0ranswer§") + assert copied.tokens == tokenizer._encode("turn 0answer§turn 1answer§") + copy_start, copy_end = len(prompt), len(prompt) + len("answer") + assert copied.flags[copy_start:copy_end] == [ + TokenFlag.EXACT | TokenFlag.ASSISTANT | TokenFlag.OUTPUT + ] * len("answer") + assert all(math.isnan(lp) for lp in copied.logprobs[copy_start:copy_end]) + assert copied.flags[copy_end] == TokenFlag.EXACT | TokenFlag.STOP + assert copied.tokens[: len(next_prompt)] == next_prompt + standalone = trajectory.histories()[1] + assert isinstance(standalone, tr.ChatCompletionsHistory) + with pytest.raises(ValueError, match="complete original sampled occurrence"): + standalone.tokenize(tokenizer=tokenizer) + + +@pytest.mark.parametrize( + "tamper", ["model", "owner", "ids", "logprob", "sampled", "trace"] +) +def test_copied_context_requires_actual_prior_source_ownership(tamper: str) -> None: + from test_tokenize import _chat_exchange + + import art.trajectories as tr + + first = _chat_exchange([1], [2, 3]) + trajectory = tr.Trajectory( + exchanges=tr.TrajectoryExchanges(chat_completions=[first]) + ) + result, traces = module._tokenize_trajectory_with_trace(trajectory) + prior, trace = result.histories[0], traces[0] + assert isinstance(prior.history, tr.ChatCompletionsHistory) + source = prior.history.message_sources[1] + assert source is not None + assert module._complete_source_is_represented( + source.model_copy(), [1], [2, 3], [-0.2, -0.3], [(prior, trace)] + ) + key = module._sampled_source_key(source) + if tamper == "model": + prior.model = "other/model" + elif tamper == "owner": + trace.sources[key] = source.model_copy( + update={"exchange": first.model_copy(deep=True)} + ) + elif tamper == "ids": + prior.tokens[0] = 999 + elif tamper == "logprob": + prior.logprobs[-1] = -999 + elif tamper == "sampled": + prior.flags[-1] &= ~TokenFlag.SAMPLED + else: + trace.source_keys[-1] = None + assert not module._complete_source_is_represented( + source, [1], [2, 3], [-0.2, -0.3], [(prior, trace)] + ) + + +def test_context_copy_survives_compact_input_and_result_roundtrip() -> None: + from test_tokenize import _chat_exchange + + import art.trajectories as tr + + trajectory = tr.Trajectory( + exchanges=tr.TrajectoryExchanges( + chat_completions=[ + _chat_exchange([1], [2, 3]), + _chat_exchange([1, 3, 4], [5], offset=1), + ] + ) + ) + restored = tr.compact_validate(tr.compact_dump(trajectory), type=tr.Trajectory) + result = trajectory.tokenize(multi_history=True) + repeated = restored.tokenize(multi_history=True) + decoded = tr.compact_validate( + tr.compact_dump(result), type=tr.TokenizedMultiHistoryTrajectory + ) + assert ( + result.model_dump_json() + == repeated.model_dump_json() + == decoded.model_dump_json() + ) + + +@pytest.mark.parametrize("change", ["template", "kwargs", "edited"]) +def test_copied_context_explicit_rendering_keeps_generic_route( + monkeypatch: pytest.MonkeyPatch, change: str +) -> None: + from test_tokenize import _chat_exchange, _FakeTokenizer + + import art.trajectories as tr + + trajectory = tr.Trajectory( + exchanges=tr.TrajectoryExchanges( + chat_completions=[ + _chat_exchange([1], [2, 3]), + _chat_exchange([1, 3, 4], [5], offset=1), + ] + ) + ) + history = trajectory.histories()[1] + assert isinstance(history, tr.ChatCompletionsHistory) + assert module._partial_native_context(history) + kwargs: dict[str, Any] = {} + if change == "template": + kwargs["chat_template"] = "explicit public renderer" + elif change == "kwargs": + kwargs["chat_template_kwargs"] = {"enable_thinking": False} + else: + history.messages[0]["content"] = "edited context" + history.message_sources[0] = None + tokenizer = _FakeTokenizer() + + def not_native(*args: Any, **kwargs: Any) -> Any: + pytest.fail("explicit/edited rendering must not require prior native ownership") + + monkeypatch.setattr(module, "_complete_source_is_represented", not_native) + monkeypatch.setattr(module, "_certify_copied_context", not_native) + + def outcome() -> Any: + try: + value = history.tokenize(tokenizer=tokenizer, **kwargs) + except (ValueError, AssertionError) as error: + return type(error), str(error) + return ( + value.tokens, + value.flags, + [None if math.isnan(x) else x for x in value.logprobs], + ) + + candidate = outcome() + assert tokenizer.calls # The real generic renderer was reached. + monkeypatch.setattr(module, "_partial_native_context", lambda history: []) + assert outcome() == candidate + + +def test_native_record_reuse_is_local_and_observes_later_mutation( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from collections import Counter + + from test_tokenize import _chat_exchange + + import art.trajectories as tr + + first = _chat_exchange([1], [2, 3]) + second = _chat_exchange([1, 2, 3, 4], [5, 6], offset=1) + first.response.choices[0].index = 7 # Choice indices are not list positions. + trajectory = tr.Trajectory( + exchanges=tr.TrajectoryExchanges(chat_completions=[first, second]) + ) + other = _chat_exchange([8], [9]) + other.request["model"] = "other/model" + other.response.model = "other/model" + nested = tr.Trajectory(exchanges=tr.TrajectoryExchanges(chat_completions=[other])) + read = module._chat_source_record + calls: Counter[int] = Counter() + nested_results = [] + + def observe(source: object) -> Any: + exchange = getattr(source, "exchange") + calls[id(exchange)] += 1 + if exchange is first: + nested_results.append(nested.tokenize()) + return read(source) + + monkeypatch.setattr(module, "_chat_source_record", observe) + result = trajectory.tokenize() + assert result.tokens == [1, 2, 3, 4, 5, 6] + assert calls[id(first)] == calls[id(second)] == calls[id(other)] == 1 + assert nested_results[0].tokens == [8, 9] + assert nested_results[0].model == "other/model" + lp = first.response.choices[0].logprobs + assert lp is not None and lp.content is not None + lp.content[1].logprob = -7.5 + repeated = trajectory.tokenize() + assert repeated.logprobs[2] == -7.5 + assert result.logprobs[2] == -0.3 + assert calls[id(first)] == calls[id(second)] == calls[id(other)] == 2 + + failure = ValueError("public native record failure") + + def fail(source: object) -> Any: + raise failure + + monkeypatch.setattr(module, "_chat_source_record", fail) + with pytest.raises(ValueError) as caught: + trajectory.tokenize() + assert caught.value is failure + monkeypatch.setattr(module, "_chat_source_record", read) + assert trajectory.tokenize().logprobs[2] == -7.5 diff --git a/tests/unit/trajectories/test_tokenize.py b/tests/unit/trajectories/test_tokenize.py index 0e04a5268..23c1452ff 100644 --- a/tests/unit/trajectories/test_tokenize.py +++ b/tests/unit/trajectories/test_tokenize.py @@ -4624,12 +4624,14 @@ def test_cross_exchange_responses_reasoning_split_uses_later_prompt_backbone() - [1, 3, 4, 5], ] assert math.isnan(tokenized.histories[1].logprobs[0]) - assert tokenized.histories[1].logprobs[1] == -0.3 + # The copied 3 was sampled after [1, 2], never after [1]. + assert tokenized.histories[0].logprobs[1:] == [-0.2, -0.3] + assert math.isnan(tokenized.histories[1].logprobs[1]) assert math.isnan(tokenized.histories[1].logprobs[2]) assert tokenized.histories[1].logprobs[3] == -0.1 assert tokenized.histories[1].flags == [ tr.TokenFlag.EXACT, - _SAMPLED_ASSISTANT_OUTPUT, + tr.TokenFlag.EXACT | tr.TokenFlag.ASSISTANT | tr.TokenFlag.OUTPUT, tr.TokenFlag.EXACT, _SAMPLED_ASSISTANT_OUTPUT, ] @@ -7215,13 +7217,19 @@ def apply_chat_template( [1, 2, 101, 102, 9], [1, 101, 102, 9, 4, 5, 6, 9], ] - assert tokenized.histories[1].flags[1] & tr.TokenFlag.SAMPLED + assert not tokenized.histories[1].flags[1] & tr.TokenFlag.SAMPLED + assert tokenized.histories[1].flags[1] & tr.TokenFlag.OUTPUT + assert tokenized.histories[0].logprobs[2:4] == [-10.1, -10.2] assert tokenized.histories[1].flags[1] & tr.TokenFlag.EXACT - assert tokenized.histories[1].logprobs[1:3] == [-10.1, -10.2] + assert all(math.isnan(value) for value in tokenized.histories[1].logprobs[1:3]) assert tokenized.histories[1].flags[3] == ( - _SAMPLED_ASSISTANT_OUTPUT | tr.TokenFlag.STOP + tr.TokenFlag.EXACT + | tr.TokenFlag.ASSISTANT + | tr.TokenFlag.OUTPUT + | tr.TokenFlag.STOP ) - assert tokenized.histories[1].logprobs[3] == -0.9 + assert math.isnan(tokenized.histories[1].logprobs[3]) + assert tokenized.histories[0].logprobs[4] == -0.9 assert 2 not in tokenized.histories[1].tokens assert 500 not in tokenized.histories[1].tokens @@ -7381,13 +7389,21 @@ def apply_chat_template( multi_history=True, tokenizer=Tokenizer(), ) - second_history = tokenized.histories[1] + first_history, second_history = tokenized.histories + assert first_history.tokens == [1, 2, 7, 8] + assert first_history.logprobs[1:] == [-0.2, -0.7, -0.8] + assert first_history.flags[1:] == [_SAMPLED_ASSISTANT_OUTPUT] * 3 assert second_history.tokens == [1, 7, 8, 4, 5] - assert second_history.logprobs[1:3] == [-0.7, -0.8] - assert second_history.flags[1:3] == [ - _SAMPLED_ASSISTANT_OUTPUT, - _SAMPLED_ASSISTANT_OUTPUT, - ] + assert all(math.isnan(lp) for lp in second_history.logprobs[1:3]) + assert ( + second_history.flags[1:3] + == [ + tr.TokenFlag.EXACT | tr.TokenFlag.ASSISTANT | tr.TokenFlag.OUTPUT, + ] + * 2 + ) + assert second_history.logprobs[-1] == -0.5 + assert second_history.flags[-1] == _SAMPLED_ASSISTANT_OUTPUT preprocessing = list( tokenize_trajectory_groups( From 09018c982a3f0e927ef486895c98df197b00f695 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 20:00:29 +0000 Subject: [PATCH 05/16] Test copied-context loss and routing after length filtering --- .../test_exchange_training_model_selection.py | 53 +++++++++++++------ 1 file changed, 38 insertions(+), 15 deletions(-) diff --git a/tests/unit/test_exchange_training_model_selection.py b/tests/unit/test_exchange_training_model_selection.py index 9d9935663..6a45b409e 100644 --- a/tests/unit/test_exchange_training_model_selection.py +++ b/tests/unit/test_exchange_training_model_selection.py @@ -2,6 +2,7 @@ from collections.abc import Mapping from datetime import datetime, timedelta +import math from pathlib import Path from types import SimpleNamespace from typing import Any, SupportsIndex, cast, overload @@ -18,7 +19,7 @@ from art.dev.model import InternalModelConfig from art.local import LocalBackend from art.openai import ART_MOE_ROUTING_METADATA_KEY -from art.preprocessing.moe_routing import MoeRouteArray +from art.preprocessing.moe_routing import MoeRouteArray, MoeRouteSegments from art.preprocessing.tokenize import ( TokenizedResult, _chat_choice_trace, @@ -361,7 +362,10 @@ def counted_public( assert public_calls == len(group.trajectories) -def test_overlength_history_does_not_claim_sources_from_fitting_history() -> None: +@pytest.mark.parametrize("max_sequence_length", [None, 5]) +def test_overlength_history_does_not_promote_copied_context( + max_sequence_length: int | None, +) -> None: results = list( tokenize_trajectory_groups( cast(PreTrainedTokenizerBase, _Tokenizer()), @@ -371,17 +375,25 @@ def test_overlength_history_does_not_claim_sources_from_fitting_history() -> Non shuffle_group_trajectories=False, drop_zero_advantage_trajectories=False, model="policy", - _max_sequence_length=5, + _max_sequence_length=max_sequence_length, ) ) long = [result for result in results if len(result.token_ids) > 5] fitting = [result for result in results if len(result.token_ids) <= 5] assert len(long) == len(fitting) == 2 - assert all(result.assistant_mask == [0] * 7 for result in long) + assert all(result.token_ids == [1, 2, 101, 102, 103, 104, 9] for result in long) + original_mask = [0] * 7 if max_sequence_length == 5 else [0, 1, 1, 1, 1, 1, 1] + assert all(result.assistant_mask == original_mask for result in long) + assert all(result.logprobs[1:] == [-0.1] * 6 for result in long) assert all(result.token_ids == [1, 9, 4, 5, 6] for result in fitting) - assert all(result.assistant_mask == [0, 1, 0, 1, 1] for result in fitting) - assert all(result.weight == pytest.approx(1 / 3) for result in results) + # Dropping the complete native occurrence cannot make token 9 sampled under + # [1]: its recorded logprob was conditioned on [1, 2, 101, 102, 103, 104]. + assert all(result.assistant_mask == [0, 0, 0, 1, 1] for result in fitting) + assert all(math.isnan(result.logprobs[1]) for result in fitting) + assert all(result.logprobs[3:] == [-0.1, -0.1] for result in fitting) + denominator = 2 if max_sequence_length == 5 else 8 + assert all(result.weight == pytest.approx(1 / denominator) for result in results) def test_local_backend_trains_retained_source_after_overlength_history( @@ -424,7 +436,7 @@ def test_local_backend_trains_retained_source_after_overlength_history( assert packed is not None assert packed["tokens"].tolist() == [[1, 9, 4, 5, 6]] * 2 - assert packed["assistant_mask"].tolist() == [[False, True, False, True, True]] * 2 + assert packed["assistant_mask"].tolist() == [[False, False, False, True, True]] * 2 def test_training_rejects_multiple_concrete_policy_versions() -> None: @@ -800,19 +812,30 @@ def apply_chat_template( assert len(initial) == 2 assert len(stripped) == 2 assert all(result.choice_offsets == [1] for result in initial) - # The retained response has a different complete visible prefix after its - # reasoning is stripped, so it is independently eligible in this history. - assert all(result.choice_offsets == [1, 5] for result in stripped) + # The copied suffix has different conditioning, so only the later complete + # response is sampled here. Its recorded prompt still supplies MoE routes. + assert all(result.choice_offsets == [5] for result in stripped) assert all(result.assistant_mask == [0, 1, 1, 1, 1] for result in initial) - assert all(result.assistant_mask == [0, 1, 1, 1, 0, 1, 1] for result in stripped) - assert all(result.weight == pytest.approx(1 / 9) for result in results) + assert all(result.logprobs[1:] == [-0.2, -10.1, -10.2, -0.9] for result in initial) + assert all(result.assistant_mask == [0, 0, 0, 0, 0, 1, 1] for result in stripped) + assert all(all(math.isnan(lp) for lp in result.logprobs[:5]) for result in stripped) + assert all(result.logprobs[5:] == [-0.5, -0.6] for result in stripped) + assert all(result.weight == pytest.approx(1 / 6) for result in results) expected_routes = np.asarray( [[[10]], [[1010]], [[1020]], [[90]], [[40]], [[50]], [[60]]], dtype=np.uint16, ) for result in stripped: - assert isinstance(result.moe_routed_experts, MoeRouteArray) - assert np.array_equal(result.moe_routed_experts, expected_routes) + assert isinstance(result.moe_routed_experts, MoeRouteSegments) + assert np.array_equal( + np.concatenate(result.moe_routed_experts.segments), expected_routes + ) + for result in initial: + assert isinstance(result.moe_routed_experts, MoeRouteSegments) + assert np.array_equal( + np.concatenate(result.moe_routed_experts.segments), + np.asarray([[[10]], [[20]], [[1010]], [[1020]], [[90]]], dtype=np.uint16), + ) datums = trajectory_groups_to_datums( [group], @@ -824,7 +847,7 @@ def apply_chat_template( ) masks = [datum.loss_fn_inputs["mask"].to_torch().tolist() for datum in datums] assert masks.count([1, 1, 1, 1]) == 2 - assert masks.count([1, 1, 1, 0, 1, 1]) == 2 + assert masks.count([0, 0, 0, 0, 1, 1]) == 2 def test_ambiguous_non_moe_suffix_falls_back_to_sampled_spans() -> None: From ab0b3ea1f647b400296c7f72b13ff2e94a797742 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 20:51:43 +0000 Subject: [PATCH 06/16] Preserve copied-context authority across native protocols --- src/art/trajectories/_tokenize.py | 109 ++++-- src/art_inference/chat_template.py | 13 +- tests/unit/test_literal_reasoning_content.py | 40 ++ .../trajectories/test_recorded_boundaries.py | 367 ++++++++++++++++++ tests/unit/trajectories/test_tokenize.py | 31 +- 5 files changed, 523 insertions(+), 37 deletions(-) diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index e13db82be..d1399dcc2 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -809,16 +809,19 @@ def _require_causal_predecessor(trainable: Sequence[bool]) -> None: @dataclass class _TraceBuilder: trace: _HistoryTokenizationTrace | None = None + rendered_outputs: tuple[tuple[int, int, object], ...] = () def set( self, tokenized: TokenizedHistory, source_keys: list[_SampledSourceKey | None], sources: dict[_SampledSourceKey, object], + rendered_outputs: tuple[tuple[int, int, object], ...] = (), ) -> None: trace = _HistoryTokenizationTrace(source_keys=source_keys, sources=sources) trace.validate(tokenized) self.trace = trace + self.rendered_outputs = rendered_outputs def _fingerprint(value: object) -> str: @@ -4180,7 +4183,7 @@ def _complete_source_is_represented( ) -> bool: """Prove ownership of the original edge before treating a copy as context.""" key = _sampled_source_key(source) - exchange = getattr(source, "exchange", None) + exchange = _source_exchange(source) required = ( TokenFlag.EXACT | TokenFlag.SAMPLED | TokenFlag.ASSISTANT | TokenFlag.OUTPUT ) @@ -4190,7 +4193,7 @@ def _complete_source_is_represented( owner = trace.sources.get(key) if ( previous.model != getattr(exchange, "model", None) - or getattr(owner, "exchange", None) is not exchange + or _source_exchange(owner) is not exchange or getattr(owner, "choice_index", None) != getattr(source, "choice_index", None) or trace.source_keys[len(prompt) : end] != [key] * len(output) @@ -4214,7 +4217,13 @@ def _complete_source_is_represented( def _source_native_record( source: object, ) -> tuple[list[int] | None, list[int] | None, list[float]]: - exchange = getattr(source, "exchange", None) + exchange = _source_exchange(source) + if isinstance(exchange, MessagesExchange) and ( + isinstance(source, MessagesExchange) + or isinstance(source, AnthropicMessageSource) + and source.request_index is None + ): + return _messages_tokens(exchange.response) if isinstance(exchange, ResponsesExchange): index = getattr(source, "generation_index", None) generations = _response_generations(exchange.response) @@ -4238,7 +4247,14 @@ def _source_native_prefix(source: object) -> tuple[list[int] | None, list[int] | if prompt is None: prompt = (exchange.response.model_extra or {}).get("prompt_token_ids") output = (choice.model_extra or {}).get("token_ids") - if isinstance(prompt, list) and isinstance(output, list): + if ( + isinstance(prompt, list) + and prompt + and isinstance(output, list) + and output + and all(type(value) is int and value >= 0 for value in prompt) + and all(type(value) is int and value >= 0 for value in output) + ): return prompt, output prompt, output, _ = _source_native_record(source) return prompt, output @@ -4283,9 +4299,27 @@ def _certify_copied_context( trace: _HistoryTokenizationTrace, copied: Sequence[object], prior: Sequence[tuple[TokenizedHistory, _HistoryTokenizationTrace]], + rendered_outputs: Sequence[tuple[int, int, object]] = (), ) -> None: """A rendered copy may keep output provenance, never its old prediction LP.""" copied_keys = {_sampled_source_key(source) for source in copied} + # Visible logprob replacements need not be sampled. Keep their output + # provenance separate from the trace's strictly sampled-token ownership. + for start, end, source in rendered_outputs: + if _sampled_source_key(source) not in copied_keys: + continue + prompt, output, logprobs = _source_native_record(source) + if ( + prompt is None + or output is None + or not _complete_source_is_represented( + source, prompt, output, logprobs, prior + ) + ): + raise ValueError( + "Copied rendered output has no complete original sampled occurrence" + ) + tokenized.logprobs[start:end] = [math.nan] * (end - start) positions: dict[_SampledSourceKey, list[int]] = {} for index, key in enumerate(trace.source_keys): if key is not None: @@ -4661,11 +4695,14 @@ def _tokenize_recorded_chat_boundaries( ): continue try: - body = decode( - output, - skip_special_tokens=False, - clean_up_tokenization_spaces=False, - ) + try: + body = decode( + output, + skip_special_tokens=False, + clean_up_tokenization_spaces=False, + ) + except ValueError: + return None generation = render(messages[:index], add_generation_prompt=True) completed = render(messages[: index + 1], add_generation_prompt=False) # Only the actually sampled body anchors the tail. Literal content, @@ -4678,11 +4715,14 @@ def _tokenize_recorded_chat_boundaries( if len(stops) != 1: return None terminator = stops[0] - trailing = decode( - tail[terminator + 1 :], - skip_special_tokens=False, - clean_up_tokenization_spaces=False, - ) + try: + trailing = decode( + tail[terminator + 1 :], + skip_special_tokens=False, + clean_up_tokenization_spaces=False, + ) + except ValueError: + return None if not isinstance(trailing, str) or trailing and not trailing.isspace(): return None following: list[int] = [] @@ -6489,6 +6529,7 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: source_keys: list[_SampledSourceKey | None] = [] sources: dict[_SampledSourceKey, object] = {} cursor = 0 + rendered_outputs: list[tuple[int, int, object]] = [] for ( start, end, @@ -6572,6 +6613,10 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: source_keys.extend([source_key] * len(replacement)) sources[source_key] = source else: + if _trace is not None: + rendered_outputs.append( + (len(token_ids), len(token_ids) + len(replacement), source) + ) token_ids.extend(replacement) logprobs.extend( replacement_logprobs @@ -6647,7 +6692,7 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: flags=flags, ) if _trace is not None: - _trace.set(tokenized, source_keys, sources) + _trace.set(tokenized, source_keys, sources, tuple(rendered_outputs)) return tokenized @@ -7009,6 +7054,7 @@ def _tokenize_history( _trace: _TraceBuilder | None = None, _prior: Sequence[tuple[TokenizedHistory, _HistoryTokenizationTrace]] = (), _projection_validated: bool = False, + _copied_context: bool = False, ) -> TokenizedHistory: if isinstance(history, LegacyHistory): if model is None: @@ -7102,7 +7148,9 @@ def _tokenize_history( ), _trace=_trace, ) - if isinstance(history, AnthropicMessagesHistory) and needs_render: + if isinstance(history, AnthropicMessagesHistory) and ( + needs_render or _copied_context + ): converted = history.as_chat_completions_history() if ( not has_length_stop @@ -7120,18 +7168,22 @@ def _tokenize_history( tokenizer=tokenizer, projection_validated=True, _trace=_trace, + _strict_sources=True, + _prior=_prior, ) ) ): return exact - return _tokenize_chat_view( - converted, - base_model=base_model, - tokenizer=tokenizer, - chat_template=chat_template, - chat_template_kwargs=chat_template_kwargs, - _trace=_trace, - ) + if needs_render: + return _tokenize_chat_view( + converted, + base_model=base_model, + tokenizer=tokenizer, + chat_template=chat_template, + chat_template_kwargs=chat_template_kwargs, + _trace=_trace, + _prior=_prior, + ) if isinstance(history, ResponsesHistory) and needs_render: return _tokenize_chat_view( history.as_chat_completions_history(), @@ -7226,13 +7278,20 @@ def tokenize_history( _trace=trace_builder, _prior=_prior, _projection_validated=_projection_validated, + _copied_context=bool(copied), ) if copied: if trace_builder is None or trace_builder.trace is None: raise ValueError( "Copied native context requires a complete tokenization source trace" ) - _certify_copied_context(tokenized, trace_builder.trace, copied, _prior) + _certify_copied_context( + tokenized, + trace_builder.trace, + copied, + _prior, + trace_builder.rendered_outputs, + ) # Internal protocol conversion is an implementation detail. The source is # always the public history view the caller asked to tokenize. if not isinstance( diff --git a/src/art_inference/chat_template.py b/src/art_inference/chat_template.py index fcf8a1d0d..349d004fa 100644 --- a/src/art_inference/chat_template.py +++ b/src/art_inference/chat_template.py @@ -49,15 +49,22 @@ def _without_inline_reasoning_parser(template: str) -> str: # Only executable block tokens may be edited. The same spelling inside a # quoted expression, raw block or comment is literal template data. + # Jinja lexes normalized newlines; map its positions to the original source. + normalized = re.sub(r"\r\n?", "\n", template) + offsets = [ + i + for i, char in enumerate(template) + if not (char == "\n" and i and template[i - 1] == "\r") + ] starts: set[int] = set() cursor = 0 try: for _, kind, value in Environment().lex(template): - start = template.find(value, cursor) - if start < 0 or template[cursor:start].strip(): + start = normalized.find(value, cursor) + if start < 0 or normalized[cursor:start].strip(): return template # Lexer normalization could not be source-joined. if kind == "block_begin": - starts.add(start) + starts.add(offsets[start]) cursor = start + len(value) except TemplateSyntaxError: return template # Leave invalid templates to their existing renderer. diff --git a/tests/unit/test_literal_reasoning_content.py b/tests/unit/test_literal_reasoning_content.py index 8561dc5fa..88624f690 100644 --- a/tests/unit/test_literal_reasoning_content.py +++ b/tests/unit/test_literal_reasoning_content.py @@ -6,6 +6,8 @@ import pytest from art_inference.chat_template import ( + _QWEN_INLINE_REASONING, + _without_inline_reasoning_parser, chat_template_with_preserved_thinking, default_chat_template_kwargs_for_template, ) @@ -265,3 +267,41 @@ def test_mixed_executable_and_literal_operations_only_changes_executable(wrapper fixed, [_USER, {"role": "assistant", "content": "headliteraltail"}], ) + + +@pytest.mark.parametrize("newline", ["\n", "\r\n", "\r"]) +@pytest.mark.parametrize("prefix", ["comment", "data"]) +def test_newline_lexing_preserves_literal_content(newline, prefix): + intro = ( + "{# public\nmultiline comment #}\n" + if prefix == "comment" + else "public\nheader\n" + ) + template = (intro + _TEMPLATE).replace("\n", newline) + content = "prefixliteralsuffix" + fixed = chat_template_with_preserved_thinking(template) + assert isinstance(fixed, str) + assert content in _render( + fixed, + [_USER, {"role": "assistant", "content": content}], + enable_thinking=False, + preserve_thinking=True, + ) + assert fixed.startswith(intro.replace("\n", newline)) + assert not _QWEN_INLINE_REASONING.search(fixed) + assert chat_template_with_preserved_thinking(fixed) == fixed + + +@pytest.mark.parametrize("newline", ["\r\n", "\r"]) +@pytest.mark.parametrize("wrapper", ["comment", "raw", "quoted"]) +def test_newline_parser_spelling_in_nonexecutable_token_unchanged(newline, wrapper): + match = _QWEN_INLINE_REASONING.search(_TEMPLATE) + assert match is not None + operation = match.group().replace("\n", newline) + if wrapper == "comment": + template = "{#" + newline + operation + newline + "#}" + elif wrapper == "raw": + template = "{% raw %}" + newline + operation + newline + "{% endraw %}" + else: + template = '{{ "' + operation + '" }}' + assert _without_inline_reasoning_parser(template) == template diff --git a/tests/unit/trajectories/test_recorded_boundaries.py b/tests/unit/trajectories/test_recorded_boundaries.py index ef5db030f..6cac07d9b 100644 --- a/tests/unit/trajectories/test_recorded_boundaries.py +++ b/tests/unit/trajectories/test_recorded_boundaries.py @@ -553,3 +553,370 @@ def fail(source: object) -> Any: assert caught.value is failure monkeypatch.setattr(module, "_chat_source_record", read) assert trajectory.tokenize().logprobs[2] == -7.5 + + +@pytest.mark.parametrize("carrier", ["empty", "encoded", "mixed"]) +def test_copied_context_preflight_uses_authoritative_token_carriers( + carrier: str, +) -> None: + from test_tokenize import _chat_exchange + + import art.trajectories as tr + + first = _chat_exchange([1], [2, 3]) + second = _chat_exchange([1, 3, 4], [5], offset=1) + first_extra = first.response.choices[0].model_extra + second_extra = second.response.choices[0].model_extra + assert first_extra is not None and second_extra is not None + first_extra["token_ids"] = ( + [] if carrier == "empty" else ["token_id:2", "token_id:3"] + ) + if carrier == "encoded": + first_extra["prompt_token_ids"] = ["token_id:1"] + second_extra["prompt_token_ids"] = [ + "token_id:1", + "token_id:3", + "token_id:4", + ] + trajectory = tr.Trajectory( + exchanges=tr.TrajectoryExchanges(chat_completions=[first, second]) + ) + before = trajectory.model_dump_json() + result = trajectory.tokenize(multi_history=True) + assert [history.tokens for history in result.histories] == [[1, 2, 3], [1, 3, 4, 5]] + assert result.histories[0].logprobs[1:] == [-0.2, -0.3] + assert math.isnan(result.histories[1].logprobs[1]) + assert not result.histories[1].flags[1] & TokenFlag.SAMPLED + assert trajectory.model_dump_json() == before + + +def test_messages_copied_context_requires_and_preserves_original_owner() -> None: + from test_tokenize import _message_exchange + + import art.trajectories as tr + + first = _message_exchange( + tr.MessagesRequest( + model="test/model", + max_tokens=16, + messages=[{"role": "user", "content": "one"}], + ), + content=[ + {"type": "thinking", "thinking": "reason", "signature": "public"}, + {"type": "text", "text": "answer"}, + ], + prompt_token_ids=[1], + token_ids=[2, 3], + logprobs=[-0.2, -0.3], + ) + second = _message_exchange( + tr.MessagesRequest( + model="test/model", + max_tokens=16, + messages=[ + {"role": "user", "content": "one"}, + {"role": "assistant", "content": "answer"}, + {"role": "user", "content": "two"}, + ], + ), + identifier="message-2", + offset=1, + content=[{"type": "text", "text": "next"}], + prompt_token_ids=[1, 3, 4], + token_ids=[5], + logprobs=[-0.5], + ) + trajectory = tr.Trajectory( + exchanges=tr.TrajectoryExchanges(messages=[first, second]) + ) + before = trajectory.model_dump_json() + + class Tokenizer: + def __call__(self, text: str, **kwargs: Any) -> list[int]: + return { + "one": [1], + "reason": [2], + "answer": [3], + "two": [4], + "next": [5], + }.get(text, [99]) + + def apply_chat_template(self, messages: Any, **kwargs: Any) -> list[int]: + return { + 1: [1], + 2: [1, 2, 3] if messages[-1].get("reasoning") else [1, 3], + 3: [1, 3, 4], + 4: [1, 3, 4, 5], + }[len(messages)] + + tokenizer = Tokenizer() + result = trajectory.tokenize(multi_history=True, tokenizer=tokenizer) + assert [history.tokens for history in result.histories] == [[1, 2, 3], [1, 3, 4, 5]] + assert result.histories[0].logprobs[1:] == [-0.2, -0.3] + assert ( + result.histories[0].flags[1:] + == [ + TokenFlag.EXACT | TokenFlag.SAMPLED | TokenFlag.ASSISTANT | TokenFlag.OUTPUT + ] + * 2 + ) + assert ( + result.histories[1].flags[1] + == TokenFlag.EXACT | TokenFlag.ASSISTANT | TokenFlag.OUTPUT + ) + assert math.isnan(result.histories[1].logprobs[1]) + assert result.histories[1].logprobs[-1] == -0.5 + assert trajectory.model_dump_json() == before + standalone = trajectory.histories()[1] + assert isinstance(standalone, tr.AnthropicMessagesHistory) + with pytest.raises(ValueError, match="complete original sampled occurrence"): + standalone.tokenize(tokenizer=tokenizer) + + +def test_unsupported_native_body_decode_preserves_generic_boundary_fallback( + monkeypatch: pytest.MonkeyPatch, +) -> None: + history, tokenizer, _ = _character_template_history() + decode = tokenizer.decode + probes = [] + + def limited_decode(tokens: list[int], **kwargs: Any) -> str: + if any(token in {7001, 7002} for token in tokens): + probes.append(tuple(tokens)) + raise ValueError("served-only public token cannot be decoded") + return decode(tokens, **kwargs) + + monkeypatch.setattr(tokenizer, "decode", limited_decode) + candidate = history.tokenize(tokenizer=tokenizer) + assert probes + monkeypatch.setattr( + module, "_tokenize_recorded_chat_boundaries", lambda *a, **k: None + ) + baseline = history.tokenize(tokenizer=tokenizer) + assert_same(candidate, baseline) + + +def test_rendered_responses_copy_clears_old_logprob_without_sampling_it() -> None: + from openai.types.responses import Response + from test_tokenize import _response_exchange + + import art.trajectories as tr + + first = _response_exchange("first", 3, prompt_token_ids=[1]) + payload = first.response.model_dump(mode="python") + text = payload["output"][0] + text["content"][0]["logprobs"] = [ + { + "token": "answer", + "bytes": list(b"answer"), + "logprob": -0.3, + "top_logprobs": [], + } + ] + payload["output"] = [ + { + "id": "public-reasoning", + "type": "reasoning", + "summary": [{"type": "summary_text", "text": "think"}], + }, + text, + ] + payload["token_generations"] = [ + { + "prompt_token_ids": [1], + "output_tokens": [ + {"token_id": 2, "logprob": -0.2}, + {"token_id": 3, "logprob": -0.3}, + ], + "output_indices": [0, 1], + } + ] + first.response = Response.model_validate(payload) + second = _response_exchange( + "second", 5, previous_response_id="first", offset=1, prompt_token_ids=[1, 3, 4] + ) + payload = second.response.model_dump(mode="python") + payload["status"] = "incomplete" + payload["incomplete_details"] = {"reason": "max_output_tokens"} + payload["output"][0]["content"][0]["text"] = "next" + second.response = Response.model_validate(payload) + trajectory = tr.Trajectory( + exchanges=tr.TrajectoryExchanges(responses=[first, second]) + ) + before = trajectory.model_dump_json() + + class Tokenizer: + def __call__(self, text: str, **kwargs: Any) -> list[int]: + return { + "turn 0": [1], + "think": [2], + "answer": [3], + "turn 1": [4], + "next": [5], + }.get(text, [99]) + + def apply_chat_template(self, messages: Any, **kwargs: Any) -> list[int]: + tokens = [] + for message in messages: + if message.get("reasoning"): + tokens += self(message["reasoning"]) + if message.get("content"): + tokens += self(message["content"]) + return tokens + + result = trajectory.tokenize(multi_history=True, tokenizer=Tokenizer()) + assert [history.tokens for history in result.histories] == [[1, 2, 3], [1, 3, 4, 5]] + assert result.histories[0].logprobs[1:] == [-0.2, -0.3] + copied = result.histories[1] + assert copied.flags[1] == TokenFlag.EXACT | TokenFlag.ASSISTANT | TokenFlag.OUTPUT + assert math.isnan(copied.logprobs[1]) + assert copied.logprobs[-1] == -0.1 + assert trajectory.model_dump_json() == before + + +@pytest.mark.parametrize("opaque", ["image", "redacted_thinking"]) +def test_complete_messages_records_do_not_require_a_chat_projection( + monkeypatch: pytest.MonkeyPatch, opaque: str +) -> None: + from test_tokenize import _message_exchange + + import art.trajectories as tr + + request = tr.MessagesRequest( + model="test/model", + max_tokens=16, + messages=[{"role": "user", "content": "question"}], + ) + if opaque == "image": + request["messages"][0]["content"] = [ + { + "type": "image", + "source": { + "type": "base64", + "media_type": "image/png", + "data": "public", + }, + } + ] + exchange = _message_exchange( + request, + prompt_token_ids=[1, 2], + token_ids=[3], + logprobs=[-0.3], + content=[{"type": "redacted_thinking", "data": "public"}] + if opaque == "redacted_thinking" + else None, + ) + trajectory = tr.Trajectory(exchanges=tr.TrajectoryExchanges(messages=[exchange])) + before = trajectory.model_dump_json() + monkeypatch.setattr( + module, + "_load_tokenizer", + lambda *_: pytest.fail("complete native record must stay offline"), + ) + result = trajectory.tokenize() + assert result.tokens == [1, 2, 3] + assert result.logprobs[-1] == -0.3 + assert result.flags == [ + TokenFlag.EXACT, + TokenFlag.EXACT, + TokenFlag.EXACT | TokenFlag.SAMPLED | TokenFlag.ASSISTANT | TokenFlag.OUTPUT, + ] + assert trajectory.model_dump_json() == before + + +def _boundary_render(tokenizer: Any) -> module._ChatRender: + def render( + selected_messages: list[dict[str, Any]], *, add_generation_prompt: bool + ) -> str: + value = tokenizer.apply_chat_template( + selected_messages, + tokenize=False, + add_generation_prompt=add_generation_prompt, + ) + assert isinstance(value, str) + return value + + return render + + +def test_optional_trailing_decode_valueerror_declines(monkeypatch): + history, tokenizer, _ = _character_template_history() + decode = tokenizer.decode + trailing = [] + + def limited(tokens, **kwargs): + if not tokens: + trailing.append(True) + raise ValueError("public empty suffix unsupported") + return decode(tokens, **kwargs) + + monkeypatch.setattr(tokenizer, "decode", limited) + result = module._tokenize_recorded_chat_boundaries( + history, + [dict(message) for message in history.messages], + tokenizer=tokenizer, + render=_boundary_render(tokenizer), + _trace=None, + ) + assert trailing and result is None + + +@pytest.mark.parametrize( + "stage", ["native_record", "render", "final_builder", "decoder_runtime"] +) +def test_other_errors_propagate_same_exception(monkeypatch, stage): + history, tokenizer, _ = _character_template_history() + error = ( + RuntimeError("public decoder failure") + if stage == "decoder_runtime" + else ValueError("public required validation failed") + ) + calls = [] + + def fail(*args, **kwargs): + calls.append(True) + raise error + + if stage == "native_record": + monkeypatch.setattr(module, "_chat_source_record", fail) + elif stage == "final_builder": + monkeypatch.setattr(module, "_tokenize_exact_projected_chat_history", fail) + elif stage == "decoder_runtime": + monkeypatch.setattr(tokenizer, "decode", fail) + render = fail if stage == "render" else _boundary_render(tokenizer) + with pytest.raises(type(error)) as caught: + module._tokenize_recorded_chat_boundaries( + history, + [dict(message) for message in history.messages], + tokenizer=tokenizer, + render=render, + _trace=None, + ) + assert calls == [True] and caught.value is error + + +def test_malformed_native_record_is_still_rejected(monkeypatch): + history, tokenizer, _ = _character_template_history() + source = history.message_sources[3] + assert source is not None and isinstance( + source.exchange, module.ChatCompletionsExchange + ) + extra = source.exchange.response.choices[0].model_extra + assert extra is not None + extra["token_ids"] = ["not-an-exact-id"] + called = [] + + def render(*args, **kwargs): + called.append(True) + raise AssertionError("should not reach rendering") + + with pytest.raises(ValueError, match="token_ids"): + module._tokenize_recorded_chat_boundaries( + history, + [dict(message) for message in history.messages], + tokenizer=tokenizer, + render=render, + _trace=None, + ) + assert not called diff --git a/tests/unit/trajectories/test_tokenize.py b/tests/unit/trajectories/test_tokenize.py index 23c1452ff..3fbb01c19 100644 --- a/tests/unit/trajectories/test_tokenize.py +++ b/tests/unit/trajectories/test_tokenize.py @@ -3698,19 +3698,32 @@ def apply_chat_template( } return by_length[len(messages)] - history = art.Trajectory( - exchanges=TrajectoryExchanges(messages=[first, second]) - ).anthropic_messages_histories()[1] - tokenized = history.tokenize(tokenizer=Tokenizer()) + trajectory = art.Trajectory(exchanges=TrajectoryExchanges(messages=[first, second])) + history = trajectory.anthropic_messages_histories()[1] + if top_level_only: + with pytest.raises(ValueError, match="complete original sampled occurrence"): + history.tokenize(tokenizer=Tokenizer()) + original, tokenized = trajectory.tokenize( + multi_history=True, tokenizer=Tokenizer() + ).histories + assert original.tokens == [10, 90, 101, 102] + assert original.logprobs[1:] == pytest.approx([-9.0, -10.1, -10.2]) + assert original.flags[1:] == [_SAMPLED_ASSISTANT_OUTPUT] * 3 + assert all(math.isnan(value) for value in tokenized.logprobs[1:3]) + assert ( + tokenized.flags[1:3] + == [tr.TokenFlag.EXACT | tr.TokenFlag.ASSISTANT | tr.TokenFlag.OUTPUT] * 2 + ) + else: + # Block-level output evidence has no native prompt. Preserve this + # existing generic rendered path rather than invent conditioning. + tokenized = history.tokenize(tokenizer=Tokenizer()) + assert tokenized.logprobs[1:3] == pytest.approx([-10.1, -10.2]) + assert tokenized.flags[1:3] == [_SAMPLED_ASSISTANT_OUTPUT] * 2 assert tokenized.tokens == [10, 101, 102, 11, 91, 201] - assert tokenized.logprobs[1:3] == pytest.approx([-10.1, -10.2]) assert tokenized.logprobs[-2] == pytest.approx(-10.0) assert tokenized.logprobs[-1] == pytest.approx(-20.1) - assert tokenized.flags[1:3] == [ - _SAMPLED_ASSISTANT_OUTPUT, - _SAMPLED_ASSISTANT_OUTPUT, - ] assert tokenized.flags[-2:] == [ _SAMPLED_ASSISTANT_OUTPUT, _SAMPLED_ASSISTANT_OUTPUT, From b11ac6b576ca3f54ef5cdd45a5f79320f49d1737 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 00:42:32 +0000 Subject: [PATCH 07/16] Reuse resolved sampled stop authority and mask alignments --- docs/features/additional-histories.mdx | 7 + src/art/trajectories/_tokenize.py | 89 ++++- .../test_resolved_stop_authority.py | 320 ++++++++++++++++++ tests/unit/trajectories/test_tokenize.py | 107 ++++++ 4 files changed, 510 insertions(+), 13 deletions(-) create mode 100644 tests/unit/trajectories/test_resolved_stop_authority.py diff --git a/docs/features/additional-histories.mdx b/docs/features/additional-histories.mdx index fd20373f4..13f8dc270 100644 --- a/docs/features/additional-histories.mdx +++ b/docs/features/additional-histories.mdx @@ -154,6 +154,13 @@ native-representation option is needed. `multi_history=True` preserves the histories selected by the trajectory, including their order and model selection. Templates still own unrecorded separators, role masks, and synthetic stop tokens. +If tokenizing another history in the same trajectory already resolves a tokenizer +for the same model, ART reuses that authority to label recorded sampled stop +tokens. This does not trigger a new tokenizer load or change rendering. Complete +native histories still work offline without a tokenizer; when neither recorded +stop metadata nor resolved tokenizer authority identifies a stop, ART leaves that +label unknown rather than guessing from the final token. + For supported Chat boundaries, ART decodes the recorded body and encodes only the unrecorded separator instead of re-tokenizing the whole conversation. It checks that the separator reproduces the next recorded prompt exactly. Edited contexts, diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index d1399dcc2..e45d48761 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -553,6 +553,7 @@ def _translate_token_mask( mask: Sequence[bool], *, tokenizer: Tokenizer | None = None, + _opcodes: list[tuple[str, int, int, int, int]] | None = None, ) -> list[bool]: """Translate a token mask across a prefix replacement without guessing.""" @@ -563,9 +564,12 @@ def _translate_token_mask( translated = [False] * len(target) mapped = [False] * len(source) decode = getattr(tokenizer, "decode", None) - for tag, start, end, target_start, target_end in SequenceMatcher( - None, source, target, autojunk=False - ).get_opcodes(): + opcodes = ( + _opcodes or SequenceMatcher(None, source, target, autojunk=False).get_opcodes() + ) + if _opcodes is not None and not _opcodes: + _opcodes.extend(opcodes) + for tag, start, end, target_start, target_end in opcodes: if tag == "equal": translated[target_start:target_end] = mask[start:end] mapped[start:end] = [True] * (end - start) @@ -809,6 +813,7 @@ def _require_causal_predecessor(trainable: Sequence[bool]) -> None: @dataclass class _TraceBuilder: trace: _HistoryTokenizationTrace | None = None + tokenizer: Tokenizer | None = None rendered_outputs: tuple[tuple[int, int, object], ...] = () def set( @@ -817,7 +822,10 @@ def set( source_keys: list[_SampledSourceKey | None], sources: dict[_SampledSourceKey, object], rendered_outputs: tuple[tuple[int, int, object], ...] = (), + *, + tokenizer: Tokenizer | None = None, ) -> None: + self.tokenizer = tokenizer trace = _HistoryTokenizationTrace(source_keys=source_keys, sources=sources) trace.validate(tokenized) self.trace = trace @@ -2752,7 +2760,7 @@ def fallback_config() -> _TokenizerConfig: flags=flags, ) if _trace is not None: - _trace.set(tokenized, source_keys, sources) + _trace.set(tokenized, source_keys, sources, tokenizer=tokenizer) return tokenized @@ -3772,7 +3780,7 @@ def _tokenize_exact_responses_history( flags=flags, ) if _trace is not None: - _trace.set(tokenized, source_keys, sources) + _trace.set(tokenized, source_keys, sources, tokenizer=tokenizer) return tokenized @@ -4620,7 +4628,7 @@ def record( flags=flags, ) if _trace is not None: - _trace.set(tokenized, source_keys, sources) + _trace.set(tokenized, source_keys, sources, tokenizer=tokenizer) return tokenized @@ -5365,22 +5373,32 @@ def source_matches_context(source: object) -> bool: canonical_assistant_mask, direct_bounds or None, ) + # These four masks translate the same pair of token sequences. + mask_opcodes: list[tuple[str, int, int, int, int]] = [] assistant_mask = _translate_token_mask( canonical_rendered, rendered, canonical_assistant_mask, tokenizer=resolved_tokenizer, + _opcodes=mask_opcodes, ) output_mask = _translate_token_mask( canonical_rendered, rendered, canonical_output_mask, tokenizer=resolved_tokenizer, + _opcodes=mask_opcodes, + ) + stop_mask = _translate_token_mask( + canonical_rendered, rendered, canonical_stop_mask, _opcodes=mask_opcodes ) - stop_mask = _translate_token_mask(canonical_rendered, rendered, canonical_stop_mask) length_stop_mask = _translate_token_mask( - canonical_rendered, rendered, canonical_length_stop_mask + canonical_rendered, + rendered, + canonical_length_stop_mask, + _opcodes=mask_opcodes, ) + mask_opcodes.clear() positions_by_first_token: dict[int, list[int]] = {} for index, token_id in enumerate(rendered): positions_by_first_token.setdefault(token_id, []).append(index) @@ -6692,7 +6710,13 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: flags=flags, ) if _trace is not None: - _trace.set(tokenized, source_keys, sources, tuple(rendered_outputs)) + _trace.set( + tokenized, + source_keys, + sources, + tuple(rendered_outputs), + tokenizer=resolved_tokenizer, + ) return tokenized @@ -6770,7 +6794,7 @@ def _tokenize_completions_token_history( flags=flags, ) if _trace is not None: - _trace.set(tokenized, source_keys, sources) + _trace.set(tokenized, source_keys, sources, tokenizer=tokenizer) return tokenized @@ -6960,7 +6984,7 @@ def resolved_tokenizer() -> Tokenizer: flags=flags, ) if _trace is not None: - _trace.set(tokenized, source_keys, sources) + _trace.set(tokenized, source_keys, sources, tokenizer=tokenizer) return tokenized @@ -7323,6 +7347,36 @@ def _materialize_trajectory( ) +def _complete_resolved_sampled_stops( + tokenized: Sequence[TokenizedHistory], builders: Sequence[_TraceBuilder | None] +) -> None: + """Reuse only authority actually resolved for this model in this call. + + Do not load a tokenizer or feed it back into renderer selection. Conflicting + tokenizer objects leave that model's unknown STOP labels unchanged. + """ + resolved: dict[str, Tokenizer | None] = {} + for value, builder in zip(tokenized, builders, strict=True): + if builder is not None and builder.tokenizer is not None: + previous = resolved.setdefault(value.model, builder.tokenizer) + if previous is not builder.tokenizer: + resolved[value.model] = None + for value, builder in zip(tokenized, builders, strict=True): + if ( + builder is not None + and builder.tokenizer is None + and builder.trace is not None + and (tokenizer := resolved.get(value.model)) is not None + ): + _mark_sampled_stops( + value.tokens, + value.flags, + builder.trace.source_keys, + builder.trace.sources, + tokenizer=tokenizer, + ) + + def tokenize_trajectory( trajectory: Trajectory, *, @@ -7356,8 +7410,10 @@ def tokenize_trajectory( track_context = len(histories) > 1 and any(context_sources) prior: list[tuple[TokenizedHistory, _HistoryTokenizationTrace]] = [] tokenized = [] + stop_builders: list[_TraceBuilder | None] = [] + collect_stops = tokenizer is None and len(histories) > 1 for history, copied in zip(histories, context_sources, strict=True): - trace = _TraceBuilder() if track_context else None + trace = _TraceBuilder() if track_context or collect_stops else None result = tokenize_history( history, model=model if isinstance(history, LegacyHistory) else history.model, @@ -7371,8 +7427,11 @@ def tokenize_trajectory( _context_sources=copied, ) tokenized.append(result) - if trace is not None and trace.trace is not None: + stop_builders.append(trace) + if track_context and trace is not None and trace.trace is not None: prior.append((result, trace.trace)) + if collect_stops: + _complete_resolved_sampled_stops(tokenized, stop_builders) if not multi_history: return _materialize_trajectory(tokenized[0], trajectory) return TokenizedMultiHistoryTrajectory( @@ -7398,6 +7457,7 @@ def _tokenize_trajectory_with_trace( histories = trajectory.histories(model=model) tokenized_histories: list[TokenizedHistory] = [] traces: list[_HistoryTokenizationTrace] = [] + builders: list[_TraceBuilder] = [] for history in histories: if isinstance(history, LegacyHistory): raise AssertionError( @@ -7419,6 +7479,9 @@ def _tokenize_trajectory_with_trace( raise AssertionError("Exchange tokenization did not produce a source trace") tokenized_histories.append(tokenized) traces.append(trace_builder.trace) + builders.append(trace_builder) + if tokenizer is None: + _complete_resolved_sampled_stops(tokenized_histories, builders) return ( TokenizedMultiHistoryTrajectory( trajectory=trajectory, diff --git a/tests/unit/trajectories/test_resolved_stop_authority.py b/tests/unit/trajectories/test_resolved_stop_authority.py new file mode 100644 index 000000000..17d15af81 --- /dev/null +++ b/tests/unit/trajectories/test_resolved_stop_authority.py @@ -0,0 +1,320 @@ +from __future__ import annotations + +import math +from typing import Any, cast + +from openai.types.chat import ChatCompletionMessageParam +import pytest +from test_tokenize import _CharacterTemplateTokenizer, _chat_exchange + +import art.trajectories as tr +from art.trajectories import _tokenize as module + + +def branch(index: int, *, length: bool, model: str = "test/model", eos: int = 9): + text = f"public question {index}" + output = _CharacterTemplateTokenizer._encode("answer") + if not length: + output.append(eos) + exchange = _chat_exchange( + _CharacterTemplateTokenizer._encode(text), output, model=model, offset=index + ) + exchange.request["messages"] = cast( + list[ChatCompletionMessageParam], [{"role": "user", "content": text}] + ) + exchange.response.choices[0].finish_reason = "length" if length else "stop" + return exchange + + +def trajectory(*exchanges): + return tr.Trajectory( + exchanges=tr.TrajectoryExchanges(chat_completions=list(exchanges)) + ) + + +def bind(monkeypatch: pytest.MonkeyPatch, tokenizers: dict[str, Any]): + loads: list[str] = [] + + def config(model: str, base_model: str | None): + assert base_model is None + return module._TokenizerConfig(model, chat_template="bound template") + + def load(config): + loads.append(config.base_model) + return tokenizers[config.base_model] + + monkeypatch.setattr(module, "_tokenizer_config", config) + monkeypatch.setattr(module, "_load_tokenizer", load) + return loads + + +def same_except_stop(left, right): + assert len(left.histories) == len(right.histories) + for a, b in zip(left.histories, right.histories, strict=True): + assert a.model == b.model and a.tokens == b.tokens + assert a.history.model_dump() == b.history.model_dump() + assert all( + x == y or math.isnan(x) and math.isnan(y) + for x, y in zip(a.logprobs, b.logprobs, strict=True) + ) + assert [f & ~tr.TokenFlag.STOP for f in a.flags] == [ + f & ~tr.TokenFlag.STOP for f in b.flags + ] + for flag in (tr.TokenFlag.SAMPLED, tr.TokenFlag.OUTPUT, tr.TokenFlag.ASSISTANT): + assert tr.first_occurrence_masks( + left.histories, where=flag + ) == tr.first_occurrence_masks(right.histories, where=flag) + + +@pytest.mark.parametrize("length_first", [False, True]) +def test_resolved_model_authority_completes_other_exact_history_stops( + monkeypatch: pytest.MonkeyPatch, length_first: bool +) -> None: + tokenizer = _CharacterTemplateTokenizer() + loads = bind(monkeypatch, {"test/model": tokenizer}) + value = trajectory( + branch(0, length=length_first), branch(1, length=not length_first) + ) + original = value.model_dump() + result = value.tokenize(multi_history=True) + assert loads == ["test/model"] + assert len(result.histories) == 2 + complete = result.histories[1 if length_first else 0] + assert complete.flags[-1] & tr.TokenFlag.SAMPLED + assert complete.flags[-1] & tr.TokenFlag.STOP + supplied = value.tokenize(tokenizer=tokenizer, multi_history=True) + same_except_stop(result, supplied) + assert [h.flags for h in result.histories] == [h.flags for h in supplied.histories] + assert value.model_dump() == original + + +def test_all_exact_histories_do_not_load_or_retain_prior_call_authority(monkeypatch): + loads = bind(monkeypatch, {"test/model": _CharacterTemplateTokenizer()}) + trajectory(branch(0, length=True), branch(1, length=False)).tokenize( + multi_history=True + ) + assert loads == ["test/model"] + loads.clear() + result = trajectory(branch(0, length=False), branch(1, length=False)).tokenize( + multi_history=True + ) + assert loads == [] + assert all(not (h.flags[-1] & tr.TokenFlag.STOP) for h in result.histories) + + +def test_resolved_authority_never_crosses_model_identity(monkeypatch): + loads = bind(monkeypatch, {"model/a": _CharacterTemplateTokenizer()}) + result = trajectory( + branch(0, length=True, model="model/a"), + branch(1, length=False, model="model/b"), + ).tokenize(multi_history=True) + assert loads == ["model/a"] + assert not result.histories[1].flags[-1] & tr.TokenFlag.STOP + + +class OtherTokenizer(_CharacterTemplateTokenizer): + eos_token_id = 8 + + @staticmethod + def _encode(text: str) -> list[int]: + return [ + 8 if value == 9 else value + for value in _CharacterTemplateTokenizer._encode(text) + ] + + def convert_tokens_to_ids(self, token: str) -> int: + return 8 if token == "§" else 0 + + def decode(self, token_ids: list[int], **kwargs: object) -> str: + return super().decode( + [9 if token == 8 else token for token in token_ids], **kwargs + ) + + +def test_each_model_uses_its_own_resolved_tokenizer(monkeypatch): + loads = bind( + monkeypatch, + {"model/a": _CharacterTemplateTokenizer(), "model/b": OtherTokenizer()}, + ) + result = trajectory( + branch(0, length=True, model="model/a"), + branch(1, length=True, model="model/b"), + branch(2, length=False, model="model/a"), + branch(3, length=False, model="model/b", eos=8), + ).tokenize(multi_history=True) + assert loads == ["model/a", "model/b"] + assert all(h.flags[-1] & tr.TokenFlag.STOP for h in result.histories) + assert [h.model for h in result.histories] == [ + "model/a", + "model/a", + "model/b", + "model/b", + ] + assert [h.tokens[-1] for h in result.histories] == [9, 9, 8, 8] + + +def test_conflicting_resolved_tokenizers_do_not_authorize_another_history(monkeypatch): + bind(monkeypatch, {}) + tokenizers = iter([_CharacterTemplateTokenizer(), OtherTokenizer()]) + monkeypatch.setattr(module, "_load_tokenizer", lambda config: next(tokenizers)) + result = trajectory( + branch(0, length=True), branch(1, length=True), branch(2, length=False) + ).tokenize(multi_history=True) + assert not result.histories[-1].flags[-1] & tr.TokenFlag.STOP + + +@pytest.mark.parametrize( + "options", + [ + {"chat_template": "caller override"}, + {"chat_template_kwargs": {"mode": "caller"}}, + ], +) +def test_stop_completion_does_not_select_or_change_render_overrides( + monkeypatch, options +): + tokenizer = _CharacterTemplateTokenizer() + rendered = [] + original_render = tokenizer.apply_chat_template + + def render(messages, **kwargs): + rendered.append(dict(kwargs)) + return original_render(messages, **kwargs) + + monkeypatch.setattr(tokenizer, "apply_chat_template", render) + bind(monkeypatch, {"test/model": tokenizer}) + value = trajectory(branch(0, length=True), branch(1, length=False)) + result = value.tokenize(multi_history=True, **options) + actual_calls = list(rendered) + rendered.clear() + monkeypatch.setattr(module, "_complete_resolved_sampled_stops", lambda *args: None) + baseline = value.tokenize(multi_history=True, **options) + assert rendered == actual_calls + same_except_stop(result, baseline) + assert [h.flags for h in result.histories] == [h.flags for h in baseline.histories] + + +def test_nested_tokenization_does_not_share_authority(monkeypatch): + tokenizer = _CharacterTemplateTokenizer() + bind(monkeypatch, {"test/model": tokenizer}) + original_render = tokenizer.apply_chat_template + nested = [] + + def render(messages, **kwargs): + if not nested: + nested.append( + trajectory(branch(7, length=False), branch(8, length=False)).tokenize( + multi_history=True + ) + ) + return original_render(messages, **kwargs) + + monkeypatch.setattr(tokenizer, "apply_chat_template", render) + result = trajectory(branch(0, length=True), branch(1, length=False)).tokenize( + multi_history=True + ) + assert result.histories[-1].flags[-1] & tr.TokenFlag.STOP + assert all(not (h.flags[-1] & tr.TokenFlag.STOP) for h in nested[0].histories) + + +def test_private_trace_and_public_results_agree(monkeypatch): + bind(monkeypatch, {"test/model": _CharacterTemplateTokenizer()}) + value = trajectory(branch(0, length=False), branch(1, length=True)) + public = value.tokenize(multi_history=True) + traced, traces = module._tokenize_trajectory_with_trace(value) + same_except_stop(public, traced) + assert [h.flags for h in public.histories] == [h.flags for h in traced.histories] + for history, trace in zip(traced.histories, traces, strict=True): + trace.validate(history) + + +def test_selected_model_does_not_load_unselected_authority(monkeypatch): + loads = bind(monkeypatch, {"model/a": _CharacterTemplateTokenizer()}) + value = trajectory( + branch(0, length=True, model="model/a"), + branch(1, length=False, model="model/b"), + ) + result = value.tokenize(model="model/b", multi_history=True) + assert loads == [] and len(result.histories) == 1 + assert not result.histories[0].flags[-1] & tr.TokenFlag.STOP + + +def test_explicit_base_is_resolved_once_without_becoming_a_render_override(monkeypatch): + calls = [] + tokenizer = _CharacterTemplateTokenizer() + + def config(model, base_model): + calls.append((model, base_model)) + return module._TokenizerConfig( + base_model, + chat_template="artifact template", + chat_template_kwargs={"public_flag": True}, + ) + + monkeypatch.setattr(module, "_tokenizer_config", config) + monkeypatch.setattr(module, "_load_tokenizer", lambda config: tokenizer) + value = trajectory(branch(0, length=False), branch(1, length=True)) + result = value.tokenize(base_model="public/base", multi_history=True) + assert calls == [("test/model", "public/base")] + assert result.histories[0].flags[-1] & tr.TokenFlag.STOP + + +@pytest.mark.parametrize("reason", [9, "§"]) +def test_recorded_stop_reason_precedence_is_preserved(monkeypatch, reason): + loads = bind(monkeypatch, {"test/model": _CharacterTemplateTokenizer()}) + complete = branch(0, length=False) + complete.response.choices[0].model_extra["stop_reason"] = reason + value = trajectory(complete, branch(1, length=True)) + result = value.tokenize(multi_history=True) + assert loads == ["test/model"] + assert result.histories[0].flags[-1] & tr.TokenFlag.STOP + + +def test_marker_encoding_failure_keeps_exception_identity(monkeypatch): + failure = RuntimeError("public stop encoder failure") + + class Tokenizer(_CharacterTemplateTokenizer): + def __call__(self, text, **kwargs): + if text == "public_stop_reason": + raise failure + return super().__call__(text, **kwargs) + + bind(monkeypatch, {"test/model": Tokenizer()}) + complete = branch(0, length=False) + complete.response.choices[0].model_extra["stop_reason"] = "public_stop_reason" + with pytest.raises(RuntimeError) as caught: + trajectory(complete, branch(1, length=True)).tokenize(multi_history=True) + assert caught.value is failure + + +def test_stop_postpass_keeps_copied_context_and_synthetic_tail_roles(monkeypatch): + bind(monkeypatch, {"test/model": _CharacterTemplateTokenizer()}) + first = _chat_exchange([1], [2, 9]) + second = _chat_exchange([1, 9, 4], [5, 9], offset=1) + value = trajectory(first, second, branch(2, length=True)) + result = value.tokenize(multi_history=True) + assert len(result.histories) == 3 + original, copied, length = result.histories + assert original.flags[-1] & tr.TokenFlag.STOP + assert copied.flags[-1] & tr.TokenFlag.STOP + assert ( + copied.flags[1] + == tr.TokenFlag.EXACT | tr.TokenFlag.ASSISTANT | tr.TokenFlag.OUTPUT + ) + assert math.isnan(copied.logprobs[1]) + assert length.flags[-1] == tr.TokenFlag.STOP + monkeypatch.setattr(module, "_complete_resolved_sampled_stops", lambda *args: None) + baseline = value.tokenize(multi_history=True) + same_except_stop(result, baseline) + assert copied.flags[1] == baseline.histories[1].flags[1] + assert length.flags == baseline.histories[2].flags + + +@pytest.mark.asyncio +async def test_public_async_default_dispatch_uses_resolved_stop_authority(monkeypatch): + loads = bind(monkeypatch, {"test/model": _CharacterTemplateTokenizer()}) + value = trajectory(branch(0, length=True), branch(1, length=False)) + results = await tr.tokenize([value], multi_history=True) + assert loads == ["test/model"] + assert len(results) == 1 and results[0].trajectory is value + assert results[0].histories[-1].flags[-1] & tr.TokenFlag.STOP diff --git a/tests/unit/trajectories/test_tokenize.py b/tests/unit/trajectories/test_tokenize.py index 3fbb01c19..20b730996 100644 --- a/tests/unit/trajectories/test_tokenize.py +++ b/tests/unit/trajectories/test_tokenize.py @@ -766,6 +766,113 @@ def test_merged_whitespace_token_inherits_mask_of_its_characters( assert _translate_token_mask([1, 2], [3], mask, tokenizer=tokenizer) == [any(mask)] +def test_chat_prefix_masks_share_one_alignment(monkeypatch: pytest.MonkeyPatch) -> None: + from difflib import SequenceMatcher + + from art.trajectories._tokenize import _translate_token_mask + + original = SequenceMatcher.get_opcodes + searches = 0 + + def get_opcodes(self): + nonlocal searches + frame = sys._getframe(1) + if frame.f_code is _translate_token_mask.__code__: + searches += 1 + return original(self) + + monkeypatch.setattr(SequenceMatcher, "get_opcodes", get_opcodes) + # Real history rendering changes the first prompt token and translates all + # four masks. Its exact outputs, logprobs and stop flags must still agree. + test_exact_output_boundaries_survive_prefix_order_drift_and_length_stop() + assert searches == 1 + + +def test_reused_mask_alignment_preserves_decoder_order() -> None: + from art.trajectories._tokenize import _translate_token_mask + + source, target = [1, 2], [3] + masks = [[False, False], [True, False], [False, True], [True, True]] + + def translate(shared: bool): + events: list[object] = [] + opcodes: list[tuple[str, int, int, int, int]] = [] + + class Tokenizer: + @property + def decode(self): + events.append("lookup") + + def decode(tokens, **kwargs): + events.append((tokens.copy(), kwargs)) + return "\n\n" + + return decode + + outputs = [ + _translate_token_mask( + source, + target, + mask, + tokenizer=cast(tr.Tokenizer, Tokenizer()), + _opcodes=opcodes if shared else None, + ) + for mask in masks + ] + return outputs, events + + cached = translate(True) + assert cached == translate(False) + assert cached[0] == [[False], [True], [True], [True]] + assert source == [1, 2] and target == [3] + assert masks == [[False, False], [True, False], [False, True], [True, True]] + + +@pytest.mark.parametrize( + "error", [ValueError("decode"), KeyboardInterrupt(), SystemExit(7)] +) +def test_reused_mask_alignment_preserves_decoder_exception( + error: BaseException, +) -> None: + from art.trajectories._tokenize import _translate_token_mask + + opcodes: list[tuple[str, int, int, int, int]] = [] + + def decode(tokens, **kwargs): + raise error + + tokenizer = cast(tr.Tokenizer, SimpleNamespace(decode=decode)) + for mask in ([False, False], [True, False]): + if any(mask): + with pytest.raises(type(error)) as caught: + _translate_token_mask( + [1, 2], [3], mask, tokenizer=tokenizer, _opcodes=opcodes + ) + assert caught.value is error + else: + assert _translate_token_mask( + [1, 2], [3], mask, tokenizer=tokenizer, _opcodes=opcodes + ) == [False] + + +def test_equal_mask_alignment_does_not_compute_opcodes( + monkeypatch: pytest.MonkeyPatch, +) -> None: + import difflib + + from art.trajectories._tokenize import _translate_token_mask + + def unexpected(*args, **kwargs): + raise AssertionError("equal tokens need no alignment") + + monkeypatch.setattr(difflib, "SequenceMatcher", unexpected) + opcodes: list[tuple[str, int, int, int, int]] = [] + mask = [True, False] + actual = _translate_token_mask([1, 2], [1, 2], mask, _opcodes=opcodes) + assert actual == mask and actual is not mask + assert opcodes == [] + + def test_exact_length_boundary_with_multiple_parts_and_prefix_drift() -> None: first = _chat_exchange([1], [2, 9]) second = _chat_exchange([1, 2, 9, 3], [4, 5], offset=1) From 2e4ec3cb5c816de82dcada912336ec8efa9a0fc8 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 01:26:10 +0000 Subject: [PATCH 08/16] Preserve literal content in selected and equivalent chat templates --- docs/features/additional-histories.mdx | 6 + src/art/trajectories/_tokenize.py | 29 +++- src/art_inference/chat_template.py | 85 +++++++--- tests/unit/test_literal_reasoning_content.py | 68 ++++++++ .../trajectories/test_literal_thinking_off.py | 157 ++++++++++++++++++ 5 files changed, 314 insertions(+), 31 deletions(-) diff --git a/docs/features/additional-histories.mdx b/docs/features/additional-histories.mdx index 13f8dc270..8f093fd34 100644 --- a/docs/features/additional-histories.mdx +++ b/docs/features/additional-histories.mdx @@ -43,6 +43,12 @@ format into `reasoning_content` and `content` before rendering. A leading tag pair alone cannot establish that format. Structured reasoning and the template's generation-prompt defaults retain their existing behavior. +For named template dictionaries, ART uses the tokenizer’s default, tool, or +explicitly selected template before applying the same correction. The correction +recognizes the known inline-content parsing operations, including equivalent +quoting and spacing; it does not reinterpret arbitrary custom template logic. +This also preserves literal content when no recorded token IDs are available. + By splitting each turn into a separate history, you can preserve these tokens for training: ```python diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index e45d48761..667dc8ca8 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -2070,6 +2070,25 @@ def _response_message( raise TypeError("Completions responses do not use chat templates") +def _resolved_chat_template( + tokenizer: Tokenizer, template: object, tools: object +) -> tuple[object, dict[str, Any]]: + # Preserve preselection defaults: resolving a named template must not + # silently change its generation mode. Explicit kwargs still override these. + configured = chat_template_with_preserved_thinking(template) + defaults = default_chat_template_kwargs_for_template(configured) + if isinstance(getattr(tokenizer, "chat_template", None), dict): + select = getattr(tokenizer, "get_chat_template", None) + if callable(select): + configured = chat_template_with_preserved_thinking( + select( + chat_template=template if isinstance(template, str) else None, + tools=tools, + ) + ) + return configured, defaults + + def _template_ids( tokenizer: Tokenizer, exchange: Exchange, @@ -2117,9 +2136,9 @@ def _template_ids( or config.chat_template or getattr(tokenizer, "chat_template", None) ) - template = chat_template_with_preserved_thinking(template) + template, defaults = _resolved_chat_template(tokenizer, template, tools) kwargs = { - **default_chat_template_kwargs_for_template(template), + **defaults, **explicit_kwargs, } result = tokenizer.apply_chat_template( @@ -5015,9 +5034,11 @@ def _tokenize_chat_view( tokenizer_template = getattr(resolved_tokenizer, "chat_template", None) if isinstance(tokenizer_template, str): template = tokenizer_template - template = chat_template_with_preserved_thinking(template) + template, defaults = _resolved_chat_template( + resolved_tokenizer, template, history.tools + ) kwargs = { - **default_chat_template_kwargs_for_template(template), + **defaults, **explicit_kwargs, } ends_with_assistant = bool(messages) and messages[-1].get("role") == "assistant" diff --git a/src/art_inference/chat_template.py b/src/art_inference/chat_template.py index 349d004fa..847f173bb 100644 --- a/src/art_inference/chat_template.py +++ b/src/art_inference/chat_template.py @@ -28,62 +28,93 @@ # These operations infer reasoning from arbitrary assistant content and can # discard everything before the last or between repeated tags. # Match the operations, not a model revision or the text of a particular answer. +_QWEN_INLINE_STATEMENTS = ( + "if '' in content", + "set reasoning_content = content.split('')[0].rstrip('\\n').split('')[-1].lstrip('\\n')", + "set content = content.split('')[-1].lstrip('\\n')", + "endif", +) _QWEN_INLINE_REASONING = re.compile( r"\s*".join( r"\{%[-+]?\s*" + re.escape(statement) + r"\s*[-+]?%\}" - for statement in ( - "if '' in content", - "set reasoning_content = content.split('')[0].rstrip('\\n').split('')[-1].lstrip('\\n')", - "set content = content.split('')[-1].lstrip('\\n')", - "endif", - ) + for statement in _QWEN_INLINE_STATEMENTS ) ) def _without_inline_reasoning_parser(template: str) -> str: - matches = list(_QWEN_INLINE_REASONING.finditer(template)) - if not matches: + if "reasoning_content" not in template or "split" not in template: return template from jinja2 import Environment, TemplateSyntaxError - # Only executable block tokens may be edited. The same spelling inside a - # quoted expression, raw block or comment is literal template data. - # Jinja lexes normalized newlines; map its positions to the original source. + # Compare parsed operations, not quote/spacing choices or a template hash. + # Only executable block tokens may be edited; quoted/raw/comment data stays. + env = Environment() + operation = env.parse( + "".join("{% " + statement + " %}" for statement in _QWEN_INLINE_STATEMENTS) + ).body normalized = re.sub(r"\r\n?", "\n", template) offsets = [ i for i, char in enumerate(template) if not (char == "\n" and i and template[i - 1] == "\r") - ] - starts: set[int] = set() + ] + [len(template)] + blocks: list[tuple[int, int, int, int]] = [] cursor = 0 + opening = None try: - for _, kind, value in Environment().lex(template): + for _, kind, value in env.lex(template): start = normalized.find(value, cursor) if start < 0 or normalized[cursor:start].strip(): return template # Lexer normalization could not be source-joined. if kind == "block_begin": - starts.add(offsets[start]) + opening = offsets[start], offsets[start + len(value)] + elif kind == "block_end" and opening is not None: + # The lexer can include whitespace following a right-trim tag. + end = start + value.index("%}") + 2 + blocks.append((*opening, offsets[start], offsets[end])) + opening = None cursor = start + len(value) except TemplateSyntaxError: return template # Leave invalid templates to their existing renderer. - edits = { - (match.start(), match.end()): "" for match in matches if match.start() in starts - } + edits: dict[tuple[int, int], str] = {} + for index, (start, _, _, _) in enumerate(blocks): + selected = blocks[index : index + 4] + if len(selected) != 4 or "split" not in template[start : selected[-1][3]]: + continue + if any( + template[left[3] : right[0]].strip() + for left, right in zip(selected, selected[1:]) + ): + continue + end = selected[-1][3] + try: + if env.parse(template[start:end]).body == operation: + edits[start, end] = "" + except TemplateSyntaxError: + continue if not edits: return template # Dropping structured reasoning must not trim the visible assistant body. - for content in ( - "render_content(message.content, true)|trim", - "(render_content(message.content, true) if preserve_thinking and message.role == 'assistant' else render_content(message.content, true)|trim)", - ): - statement = "{%- set content = " + content + " %}" - for match in re.finditer(re.escape(statement), template): - if match.start() in starts: - edits[match.span()] = ( - "{%- set content = (render_content(message.content, true) if message.role == 'assistant' else render_content(message.content, true)|trim) %}" + trims = [ + env.parse("{% set content = " + content + " %}").body + for content in ( + "render_content(message.content, true)|trim", + "(render_content(message.content, true) if preserve_thinking and message.role == 'assistant' else render_content(message.content, true)|trim)", + ) + ] + for start, body_start, body_end, end in blocks: + if "render_content" not in template[body_start:body_end]: + continue + try: + if env.parse(template[start:end]).body in trims: + edits[start, end] = ( + template[start:body_start] + + " set content = (render_content(message.content, true) if message.role == 'assistant' else render_content(message.content, true)|trim) " + + template[body_end:end] ) + except TemplateSyntaxError: + continue for (start, end), replacement in sorted(edits.items(), reverse=True): template = template[:start] + replacement + template[end:] return template diff --git a/tests/unit/test_literal_reasoning_content.py b/tests/unit/test_literal_reasoning_content.py index 88624f690..deb10eeef 100644 --- a/tests/unit/test_literal_reasoning_content.py +++ b/tests/unit/test_literal_reasoning_content.py @@ -305,3 +305,71 @@ def test_newline_parser_spelling_in_nonexecutable_token_unchanged(newline, wrapp else: template = '{{ "' + operation + '" }}' assert _without_inline_reasoning_parser(template) == template + + +@pytest.mark.parametrize("spelling", ["double_quotes", "spacing", "parentheses"]) +def test_equivalent_inline_operations_preserve_literal_and_structured_fields(spelling): + match = _QWEN_INLINE_REASONING.search(_TEMPLATE) + assert match is not None + operation = match.group() + if spelling == "double_quotes": + operation = operation.replace("'", '"') + elif spelling == "spacing": + operation = operation.replace("content.split", "content . split").replace( + "[0]", "[ 0 ]" + ) + else: + operation = operation.replace( + "if '' in content", "if ('' in content)" + ) + template = _TEMPLATE[: match.start()] + operation + _TEMPLATE[match.end() :] + assert template != _TEMPLATE + fixed = chat_template_with_preserved_thinking(template) + for content in _LITERALS: + for reasoning in (None, "", "explicit structured reasoning\n"): + messages = [_USER, {"role": "assistant", "content": content}] + if reasoning is not None: + messages[-1]["reasoning_content"] = reasoning + for preserve in (False, True): + kwargs = dict(enable_thinking=False, preserve_thinking=preserve) + assert _render(fixed, messages, **kwargs) == _render( + _FIXED, messages, **kwargs + ) + assert chat_template_with_preserved_thinking(fixed) == fixed + + +@pytest.mark.parametrize("wrapper", ["raw", "comment", "quoted"]) +def test_equivalent_operation_as_literal_data_is_not_edited(wrapper): + match = _QWEN_INLINE_REASONING.search(_TEMPLATE) + assert match is not None + operation = match.group().replace("'", '"') + if wrapper == "raw": + literal = "{% raw %}" + operation + "{% endraw %}" + elif wrapper == "comment": + literal = "{#" + operation + "#}" + else: + literal = "{{ '" + operation + "' }}" + assert _without_inline_reasoning_parser(literal) == literal + assert isinstance(_FIXED, str) + assert _without_inline_reasoning_parser(_TEMPLATE + literal) == _FIXED + literal + + +@pytest.mark.parametrize("change", ["different_split", "side_effect", "different_gate"]) +def test_distinct_custom_content_operations_are_not_inferred(change): + match = _QWEN_INLINE_REASONING.search(_TEMPLATE) + assert match is not None + operation = match.group() + if change == "different_split": + operation = operation.replace( + "content.split('')[-1]", "content.split('')[0]" + ) + elif change == "side_effect": + operation = operation.replace( + "{%- endif %}", "{%- set other = content %}{%- endif %}" + ) + else: + operation = operation.replace( + "if '' in content", "if custom and '' in content" + ) + assert operation != match.group() + assert _without_inline_reasoning_parser(operation) == operation diff --git a/tests/unit/trajectories/test_literal_thinking_off.py b/tests/unit/trajectories/test_literal_thinking_off.py index ca423cd41..bd2459830 100644 --- a/tests/unit/trajectories/test_literal_thinking_off.py +++ b/tests/unit/trajectories/test_literal_thinking_off.py @@ -12,6 +12,7 @@ import art.trajectories as tr from art.trajectories import _tokenize +from art_inference.chat_template import chat_template_with_preserved_thinking # Public Qwen3.5 template after ART's existing thinking-preservation rewrite. _TEMPLATE = ( @@ -442,3 +443,159 @@ def observe(*args: Any, **kwargs: Any) -> tr.TokenizedHistory | None: for flag in value.flags[start:stop] ) assert history.model_dump(mode="python") == original + + +class _NamedTemplateTokenizer(_TemplateTokenizer): + chat_template: Any + + def __init__(self) -> None: + super().__init__() + self.chat_template = { + "default": _TEMPLATE, + "tool_use": _TEMPLATE + "TOOL_TEMPLATE", + "named": _TEMPLATE + "NAMED_TEMPLATE", + } + self.selected: list[str] = [] + self.settings: list[dict[str, Any]] = [] + + def get_chat_template(self, chat_template=None, tools=None): + # Transformers' named/default/tool selection contract, before rendering. + templates = self.chat_template + if chat_template is not None: + return templates.get(chat_template, chat_template) + if tools is not None and "tool_use" in templates: + return templates["tool_use"] + if "default" in templates: + return templates["default"] + raise ValueError("No default template") + + def apply_chat_template(self, messages, **kwargs): + template = self.get_chat_template( + kwargs.pop("chat_template", None), kwargs.get("tools") + ) + self.selected.append(template) + self.settings.append(deepcopy(kwargs)) + return super().apply_chat_template(messages, chat_template=template, **kwargs) + + +@pytest.mark.parametrize("selection", ["default", "tools", "named", "literal_override"]) +@pytest.mark.parametrize("route", ["history", "exchange"]) +def test_unconfigured_named_template_preserves_unrecorded_literal_content( + selection, route +): + tokenizer = _NamedTemplateTokenizer() + templates_before = deepcopy(tokenizer.chat_template) + history, _ = _history() + source = history.message_sources[-1] + assert source is not None + exchange = source.exchange.model_copy(deep=True) + assert isinstance(exchange, tr.ChatCompletionsExchange) + exchange.request.pop("chat_template", None) + override = ( + "named" + if selection == "named" + else _TEMPLATE + "OVERRIDE_TEMPLATE" + if selection == "literal_override" + else None + ) + tools: list[Any] | None = ( + [ + { + "type": "function", + "function": {"name": "lookup", "parameters": {"type": "object"}}, + } + ] + if selection == "tools" + else None + ) + if route == "history": + # No native tokens: preservation must come from rendering itself. + history = tr.ChatCompletionsHistory( + model="public/qwen35", + messages=deepcopy(history.messages), + message_sources=[None] * len(history.messages), + tools=tools, + ) + before = history.model_dump() + result = history.tokenize( + tokenizer=tokenizer, + chat_template=override, + chat_template_kwargs={"enable_thinking": True, "preserve_thinking": False}, + ) + assert _LITERAL in tokenizer.decode(result.tokens) + assert not any(flag & tr.TokenFlag.SAMPLED for flag in result.flags) + assert history.model_dump() == before + else: + if tools is not None: + exchange.request["tools"] = tools + before = exchange.model_dump() + result = _tokenize._template_ids( + tokenizer, + exchange, + completed=True, + config=_tokenize._TokenizerConfig(base_model="public/qwen35"), + chat_template=override, + chat_template_kwargs={"enable_thinking": True, "preserve_thinking": False}, + ) + assert _LITERAL in tokenizer.decode(result) + assert exchange.model_dump() == before + assert tokenizer.chat_template == templates_before + assert tokenizer.selected + expected = ( + templates_before["named"] + if selection == "named" + else override + if selection == "literal_override" + else templates_before["tool_use"] + if selection == "tools" + else templates_before["default"] + ) + assert all( + selected == chat_template_with_preserved_thinking(expected) + for selected in tokenizer.selected + ) + assert all( + settings["enable_thinking"] is True and settings["preserve_thinking"] is False + for settings in tokenizer.settings + ) + + +def test_named_template_selection_failure_keeps_original_error(): + tokenizer = _NamedTemplateTokenizer() + del tokenizer.chat_template["default"] + with pytest.raises(ValueError, match="No default template"): + _tokenize._resolved_chat_template(tokenizer, None, None) + # Explicit unrelated templates are selected unchanged, not rewritten merely + # because this tokenizer also has a known Qwen template in its dictionary. + custom = "{% for message in messages %}{{ message.content }}{% endfor %}" + assert _tokenize._resolved_chat_template(tokenizer, custom, None) == (custom, {}) + + +@pytest.mark.parametrize("selection", [None, "named"]) +def test_named_selection_preserves_implicit_generation_mode(selection): + tokenizer = _NamedTemplateTokenizer() + history, _ = _history() + source = history.message_sources[-1] + assert source is not None + exchange = source.exchange.model_copy(deep=True) + assert isinstance(exchange, tr.ChatCompletionsExchange) + exchange.request.pop("chat_template", None) + exchange.request.pop("chat_template_kwargs", None) + expected = tokenizer.apply_chat_template( + exchange.request["messages"], + chat_template=selection, + tokenize=True, + add_generation_prompt=True, + ) + tokenizer.settings.clear() + actual = _tokenize._template_ids( + tokenizer, + exchange, + completed=False, + config=_tokenize._TokenizerConfig(base_model="public/qwen35"), + chat_template=selection, + chat_template_kwargs=None, + ) + assert actual == expected + assert "enable_thinking" not in tokenizer.settings[-1] + assert "preserve_thinking" not in tokenizer.settings[-1] From 3ed85fad269beeab6e2cbca9f172402cbf843785 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 01:30:17 +0000 Subject: [PATCH 09/16] Retain whitespace coverage and safe named-template selection --- docs/features/additional-histories.mdx | 2 ++ src/art/trajectories/_tokenize.py | 23 ++++++++++---- src/art_inference/chat_template.py | 22 +++++++++++--- tests/unit/test_literal_reasoning_content.py | 17 +++++++++++ .../trajectories/test_literal_thinking_off.py | 30 +++++++++++++++++++ 5 files changed, 85 insertions(+), 9 deletions(-) diff --git a/docs/features/additional-histories.mdx b/docs/features/additional-histories.mdx index 8f093fd34..047a10bd4 100644 --- a/docs/features/additional-histories.mdx +++ b/docs/features/additional-histories.mdx @@ -48,6 +48,8 @@ explicitly selected template before applying the same correction. The correction recognizes the known inline-content parsing operations, including equivalent quoting and spacing; it does not reinterpret arbitrary custom template logic. This also preserves literal content when no recorded token IDs are available. +If a corrected template body is itself another dictionary entry’s name, ART +refuses that ambiguous selection rather than rendering a different template. By splitting each turn into a separate history, you can preserve these tokens for training: diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index 667dc8ca8..2a1644cd8 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -2080,12 +2080,25 @@ def _resolved_chat_template( if isinstance(getattr(tokenizer, "chat_template", None), dict): select = getattr(tokenizer, "get_chat_template", None) if callable(select): - configured = chat_template_with_preserved_thinking( - select( - chat_template=template if isinstance(template, str) else None, - tools=tools, - ) + selected = select( + chat_template=template if isinstance(template, str) else None, + tools=tools, ) + configured = chat_template_with_preserved_thinking(selected) + if configured == selected: + # apply_chat_template resolves names itself. Forwarding an + # unchanged body could accidentally select a second named entry. + return template, defaults + templates = getattr(tokenizer, "chat_template", None) + if ( + isinstance(configured, str) + and isinstance(templates, dict) + and configured in templates + ): + raise ValueError( + "The normalized chat template is also a template name; " + "cannot preserve the selected renderer without ambiguity" + ) return configured, defaults diff --git a/src/art_inference/chat_template.py b/src/art_inference/chat_template.py index 847f173bb..bf54218ee 100644 --- a/src/art_inference/chat_template.py +++ b/src/art_inference/chat_template.py @@ -45,14 +45,28 @@ def _without_inline_reasoning_parser(template: str) -> str: if "reasoning_content" not in template or "split" not in template: return template - from jinja2 import Environment, TemplateSyntaxError + from jinja2 import Environment, TemplateSyntaxError, nodes + from jinja2.visitor import NodeTransformer # Compare parsed operations, not quote/spacing choices or a template hash. # Only executable block tokens may be edited; quoted/raw/comment data stays. + class WithoutWhitespace(NodeTransformer): + def visit_Output(self, node: nodes.Output, *args: Any, **kwargs: Any): + if all( + isinstance(child, nodes.TemplateData) and not child.data.strip() + for child in node.nodes + ): + return None + return node + env = Environment() - operation = env.parse( + + def operations(text: str): + return WithoutWhitespace().visit(env.parse(text)).body + + operation = operations( "".join("{% " + statement + " %}" for statement in _QWEN_INLINE_STATEMENTS) - ).body + ) normalized = re.sub(r"\r\n?", "\n", template) offsets = [ i @@ -89,7 +103,7 @@ def _without_inline_reasoning_parser(template: str) -> str: continue end = selected[-1][3] try: - if env.parse(template[start:end]).body == operation: + if operations(template[start:end]) == operation: edits[start, end] = "" except TemplateSyntaxError: continue diff --git a/tests/unit/test_literal_reasoning_content.py b/tests/unit/test_literal_reasoning_content.py index deb10eeef..dff7da9f2 100644 --- a/tests/unit/test_literal_reasoning_content.py +++ b/tests/unit/test_literal_reasoning_content.py @@ -373,3 +373,20 @@ def test_distinct_custom_content_operations_are_not_inferred(change): ) assert operation != match.group() assert _without_inline_reasoning_parser(operation) == operation + + +@pytest.mark.parametrize("separator", ["\n", " ", "\r\n"]) +@pytest.mark.parametrize("quoted", [False, True]) +def test_plain_block_whitespace_keeps_prior_inline_parser_coverage(separator, quoted): + match = _QWEN_INLINE_REASONING.search(_TEMPLATE) + assert match is not None + # The prior regex admitted whitespace-only separators without trim dashes. + operation = match.group().replace("{%-", "{%").replace("-%}", "%}") + operation = operation.replace("\n", separator) + if quoted: + operation = operation.replace("'", '"') + template = _TEMPLATE[: match.start()] + operation + _TEMPLATE[match.end() :] + fixed = chat_template_with_preserved_thinking(template) + content = "HEADliteralTAIL" + assert content in _render(fixed, [_USER, {"role": "assistant", "content": content}]) + assert chat_template_with_preserved_thinking(fixed) == fixed diff --git a/tests/unit/trajectories/test_literal_thinking_off.py b/tests/unit/trajectories/test_literal_thinking_off.py index bd2459830..cabfce4b6 100644 --- a/tests/unit/trajectories/test_literal_thinking_off.py +++ b/tests/unit/trajectories/test_literal_thinking_off.py @@ -599,3 +599,33 @@ def test_named_selection_preserves_implicit_generation_mode(selection): assert actual == expected assert "enable_thinking" not in tokenizer.settings[-1] assert "preserve_thinking" not in tokenizer.settings[-1] + + +def test_unchanged_selected_body_is_not_selected_again_as_a_name(): + tokenizer = _NamedTemplateTokenizer() + tokenizer.chat_template = {"default": "named", "named": "DIFFERENT"} + history = tr.ChatCompletionsHistory( + model="public/qwen35", + messages=[{"role": "user", "content": "question"}], + message_sources=[None], + ) + before = deepcopy(tokenizer.chat_template) + assert tokenizer.decode(history.tokenize(tokenizer=tokenizer).tokens) == "named" + assert tokenizer.chat_template == before + + +def test_changed_body_colliding_with_a_name_refuses_before_wrong_renderer(): + tokenizer = _NamedTemplateTokenizer() + normalized = chat_template_with_preserved_thinking(_TEMPLATE) + assert isinstance(normalized, str) and normalized != _TEMPLATE + tokenizer.chat_template = {"default": _TEMPLATE, normalized: "DIFFERENT"} + before = deepcopy(tokenizer.chat_template) + history = tr.ChatCompletionsHistory( + model="public/qwen35", + messages=[{"role": "assistant", "content": _LITERAL}], + message_sources=[None], + ) + with pytest.raises(ValueError, match="also a template name"): + history.tokenize(tokenizer=tokenizer) + assert not tokenizer.calls + assert tokenizer.chat_template == before From 6c595a9a14ebd72832a5e0161cc1ce6e6c7a3ec8 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 03:21:23 +0000 Subject: [PATCH 10/16] Preserve recorded request role boundaries and reuse scoped evidence --- docs/features/additional-histories.mdx | 7 + src/art/trajectories/_tokenize.py | 318 +++++++++++-- .../unit/trajectories/test_evidence_reuse.py | 257 +++++++++++ .../test_recorded_prompt_roles.py | 432 ++++++++++++++++++ 4 files changed, 974 insertions(+), 40 deletions(-) create mode 100644 tests/unit/trajectories/test_evidence_reuse.py create mode 100644 tests/unit/trajectories/test_recorded_prompt_roles.py diff --git a/docs/features/additional-histories.mdx b/docs/features/additional-histories.mdx index 047a10bd4..433f6ec68 100644 --- a/docs/features/additional-histories.mdx +++ b/docs/features/additional-histories.mdx @@ -175,6 +175,13 @@ that the separator reproduces the next recorded prompt exactly. Edited contexts, explicit template overrides, incomplete projections, and unsupported templates continue through the generic rendering path and its source validation. +Correcting literal-content rendering does not rewrite a recorded request. When +the original request and template reproduce its complete native prompt, ART can +recover historical assistant roles from that rendering. This uses the original +tool serialization order and preserves role labels through exact length-stop +assembly; it does not restore destructive parsing for new response content. +Unproved historical role mappings still use the strict existing fallback. + A response copied into a later, shortened prompt is output provenance, but it is not a fresh sample under that new prompt. ART keeps its `OUTPUT`, `ASSISTANT`, `EXACT`, and proven `STOP` flags while removing `SAMPLED` and the old conditional diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index 2a1644cd8..1ae9cdc8f 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -595,6 +595,106 @@ def _translate_token_mask( return translated +def _recorded_prompt_role_masks( + messages: list[dict[str, Any]], + sources: Sequence[object | None], + prompt: list[int], + *, + tokenizer: Tokenizer, + template: object, + tools: object, + kwargs: Mapping[str, object], +) -> tuple[list[bool], list[bool]] | None: + """Prove roles in recorded request context with its original renderer. + + Normalizing a template corrects future literal-content rendering; it does + not change an already served prompt. Only request context is admitted here: + sampled outputs retain their separate exact conditioning/ownership proofs. + """ + if len(messages) != len(sources) or any( + source is not None and _source_is_sampled(source) for source in sources + ): + return None + if not any(message.get("role") == "assistant" for message in messages): + return None + messages, tools, kwargs = deepcopy((messages, tools, dict(kwargs))) + original_context = _render_context_key([messages, tools, dict(kwargs)]) + + def check_context() -> None: + if _render_context_key([messages, tools, dict(kwargs)]) != original_context: + raise ValueError( + "Renderer changed context while proving recorded request roles" + ) + + def render(selected: list[dict[str, Any]], *, add_generation_prompt: bool) -> str: + value = tokenizer.apply_chat_template( + normalize_tool_call_arguments_for_chat_template(selected, template), + tools=tools, + tokenize=False, + add_generation_prompt=add_generation_prompt, + **({"chat_template": template} if template is not None else {}), + **kwargs, + ) + check_context() + if not isinstance(value, str): + raise TypeError("Historical chat template did not render text") + return value + + rendered_prompt = render(messages, add_generation_prompt=True) + encoded = cast(_OffsetTokenizer, tokenizer)( + rendered_prompt, add_special_tokens=False, return_offsets_mapping=True + ) + check_context() + if _ids(encoded) != prompt: + return None + offsets = _field(encoded, "offset_mapping") + if not isinstance(offsets, list) or len(offsets) != len(prompt): + return None + characters = [False] * len(rendered_prompt) + previous_end = 0 + for index, message in enumerate(messages): + if message.get("role") != "assistant": + continue + prior = render(messages[:index], add_generation_prompt=False) + generation = render(messages[:index], add_generation_prompt=True) + completed = render(messages[: index + 1], add_generation_prompt=False) + start = _common_prefix_length(generation, completed) + if ( + start < len(prior) + or generation[: len(prior)] != prior + or completed[: len(prior)] != prior + or rendered_prompt[: len(completed)] != completed + or start < previous_end + ): + return None + characters[start : len(completed)] = [True] * (len(completed) - start) + previous_end = len(completed) + assistant = [] + previous_start = 0 + for offset in offsets: + if ( + not isinstance(offset, (list, tuple)) + or len(offset) != 2 + or any(type(value) is not int for value in offset) + ): + return None + start, end = cast(tuple[int, int], offset) + if not previous_start <= start < end <= len(characters): + return None + selected = characters[start:end] + if ( + any(selected) + and not all(selected) + and not rendered_prompt[start:end].isspace() + ): + return None + assistant.append(any(selected)) + previous_start = start + masks = _assistant_stop_masks(prompt, assistant, tokenizer) + check_context() + return masks + + def _prove_exact_sampled_assistant_span( matches: Sequence[tuple[int, int]], assistant_mask: Sequence[bool], @@ -870,7 +970,17 @@ def _sampled_evidence_fingerprint( *, protocol: Literal["chat_completions", "responses", "messages", "completions"], index: int, + _cache: dict[tuple[int, str, int], tuple[Exchange, str]] | None = None, ) -> str: + if _cache is not None: + key = (id(exchange), protocol, index) + cached = _cache.get(key) + if cached is not None and cached[0] is exchange: + return cached[1] + value = _sampled_evidence_fingerprint(exchange, protocol=protocol, index=index) + if len(_cache) < 256: + _cache[key] = (exchange, value) + return value if protocol == "chat_completions": if not isinstance(exchange, ChatCompletionsExchange): raise TypeError("Chat source has the wrong exchange type") @@ -967,6 +1077,7 @@ def _source_key( protocol: Literal["chat_completions", "responses", "messages", "completions"], index: int, prompt_index: int | None = None, + _fingerprints: dict[tuple[int, str, int], tuple[Exchange, str]] | None = None, ) -> _SampledSourceKey: return _SampledSourceKey( protocol=protocol, @@ -981,27 +1092,40 @@ def _source_key( # IDs remain part of the identity: equal output evidence can have # different causal contexts even when response IDs are reused. evidence_fingerprint=_sampled_evidence_fingerprint( - exchange, protocol=protocol, index=index + exchange, protocol=protocol, index=index, _cache=_fingerprints ), ) -def _sampled_source_key(source: object) -> _SampledSourceKey: +def _sampled_source_key( + source: object, + *, + _fingerprints: dict[tuple[int, str, int], tuple[Exchange, str]] | None = None, +) -> _SampledSourceKey: exchange = getattr(source, "exchange", None) if isinstance(exchange, ChatCompletionsExchange): index = getattr(source, "choice_index", None) if not isinstance(index, int) or isinstance(index, bool): raise ValueError("Sampled Chat source has no choice index") - return _source_key(exchange, protocol="chat_completions", index=index) + return _source_key( + exchange, + protocol="chat_completions", + index=index, + _fingerprints=_fingerprints, + ) if isinstance(exchange, ResponsesExchange): index = getattr(source, "generation_index", None) if index is None and not _response_generations(exchange.response): index = 0 if not isinstance(index, int) or isinstance(index, bool): raise ValueError("Sampled Responses source has no generation identity") - return _source_key(exchange, protocol="responses", index=index) + return _source_key( + exchange, protocol="responses", index=index, _fingerprints=_fingerprints + ) if isinstance(exchange, MessagesExchange): - return _source_key(exchange, protocol="messages", index=0) + return _source_key( + exchange, protocol="messages", index=0, _fingerprints=_fingerprints + ) if isinstance(exchange, CompletionsExchange): index = getattr(source, "choice_index", None) prompt_index = getattr(source, "prompt_index", None) @@ -1014,6 +1138,7 @@ def _sampled_source_key(source: object) -> _SampledSourceKey: protocol="completions", index=index, prompt_index=prompt_index, + _fingerprints=_fingerprints, ) raise ValueError("Sampled token source has an unsupported exchange") @@ -3442,7 +3567,11 @@ def _history_render_state(history: History) -> _HistoryRenderState: return _HistoryRenderState(needs_render=False) -def _source_signature(source: object) -> tuple[object, ...] | None: +def _source_signature( + source: object, + *, + _fingerprints: dict[tuple[int, str, int], tuple[Exchange, str]] | None = None, +) -> tuple[object, ...] | None: if source is None: return None exchange = getattr(source, "exchange", None) @@ -3462,18 +3591,21 @@ def _source_signature(source: object) -> tuple[object, ...] | None: evidence_fingerprint: str | None = None if isinstance(exchange, ChatCompletionsExchange) and isinstance(choice_index, int): evidence_fingerprint = _sampled_evidence_fingerprint( - exchange, protocol="chat_completions", index=choice_index + exchange, + protocol="chat_completions", + index=choice_index, + _cache=_fingerprints, ) elif isinstance(exchange, ResponsesExchange) and isinstance(generation_index, int): evidence_fingerprint = _sampled_evidence_fingerprint( - exchange, protocol="responses", index=generation_index + exchange, protocol="responses", index=generation_index, _cache=_fingerprints ) elif isinstance(exchange, MessagesExchange) and ( getattr(source, "output_index", None) == 0 or _chat_output_indices(source) == (0,) ): evidence_fingerprint = _sampled_evidence_fingerprint( - exchange, protocol="messages", index=0 + exchange, protocol="messages", index=0, _cache=_fingerprints ) return ( type(source), @@ -3492,8 +3624,9 @@ def _source_signature(source: object) -> tuple[object, ...] | None: def _sources_match(left: Sequence[object], right: Sequence[object]) -> bool: - return [_source_signature(item) for item in left] == [ - _source_signature(item) for item in right + fingerprints: dict[tuple[int, str, int], tuple[Exchange, str]] = {} + return [_source_signature(item, _fingerprints=fingerprints) for item in left] == [ + _source_signature(item, _fingerprints=fingerprints) for item in right ] @@ -4403,10 +4536,12 @@ def _tokenize_exact_projected_chat_history( ) -> TokenizedHistory | None: if not projection_validated and not _history_matches_projection(history): return None + # Reuse evidence only in this callback-free phase, never across render/decode. + fingerprints: dict[tuple[int, str, int], tuple[Exchange, str]] = {} sampled_sources: list[object] = [] seen: set[tuple[object, ...]] = set() for message, source in zip(history.messages, history.message_sources, strict=True): - signature = _source_signature(source) + signature = _source_signature(source, _fingerprints=fingerprints) if ( message.get("role") == "assistant" and source is not None @@ -4434,7 +4569,7 @@ def record( final_prompt, final_output, final_logprobs = record(final_source) if final_prompt is None or final_output is None: return None - final_key = _sampled_source_key(final_source) + final_key = _sampled_source_key(final_source, _fingerprints=fingerprints) final_stop_reason = _source_stop_evidence(final_source, final_key)[0] # A terminal synthetic stop can accompany an earlier length-stop boundary; # neither tail is sampled, and both must retain their renderer proof. @@ -4542,7 +4677,7 @@ def record( TokenFlag.EXACT | TokenFlag.SAMPLED | TokenFlag.ASSISTANT | TokenFlag.OUTPUT ] * len(retained_ids) logprobs[start:end] = retained_logprobs - source_key = _sampled_source_key(source) + source_key = _sampled_source_key(source, _fingerprints=fingerprints) if _strict_sources and retained_ids != output: if not _complete_source_is_represented( source, prompt, output, output_logprobs, _prior @@ -4557,6 +4692,7 @@ def record( ] * len(retained_ids) logprobs[start:end] = [math.nan] * len(retained_ids) records.clear() # A custom STOP decoder may change source objects. + fingerprints.clear() stop_count = _sampled_stop_suffix( output, source=source, source_key=source_key, tokenizer=tokenizer ) @@ -4603,6 +4739,7 @@ def record( and callable(decode) ): records.clear() # Never reuse records across user callbacks. + fingerprints.clear() if decode(native_boundary[:extra]).isspace(): # Services may insert whitespace before a truncated turn's # proven stop tail. Keep those served, nonsampled tokens. @@ -5047,6 +5184,7 @@ def _tokenize_chat_view( tokenizer_template = getattr(resolved_tokenizer, "chat_template", None) if isinstance(tokenizer_template, str): template = tokenizer_template + original_template = template template, defaults = _resolved_chat_template( resolved_tokenizer, template, history.tools ) @@ -5371,6 +5509,7 @@ def source_matches_context(source: object) -> bool: canonical_rendered = rendered exact_prefix_length = 0 canonical_prefix_length = 0 + recorded_prompt_masks = None if ( chat_template is None and chat_template_kwargs is None @@ -5393,6 +5532,72 @@ def source_matches_context(source: object) -> bool: rendered = [*source_prompt, *rendered[len(rendered_prompt) :]] exact_prefix_length = len(source_prompt) canonical_prefix_length = len(rendered_prompt) + if ( + original_template != template + and source_prompt != rendered_prompt + ): + signature = _source_signature(source) + request_context = None + try: + # Canonical tool validation may reorder JSON keys. + # Prove the historical prompt with the recorded + # request's own serialization inputs, not a view. + exchange = getattr(source, "exchange", None) + assert isinstance( + exchange, + ( + ChatCompletionsExchange, + MessagesExchange, + ResponsesExchange, + ), + ) + request_messages, request_tools = _request_messages( + exchange + ) + request_context = _render_context_key( + [request_messages, request_tools] + ) + recorded_prompt_masks = _recorded_prompt_role_masks( + request_messages, + history.message_sources[:message_index], + source_prompt, + tokenizer=resolved_tokenizer, + template=original_template, + tools=request_tools, + kwargs=kwargs, + ) + except (TypeError, KeyError, NotImplementedError): + # Optional historical prefix rendering may be + # unsupported although the selected renderer works. + recorded_prompt_masks = None + # These optional renderer calls must not turn cached + # native evidence into authority for a changed source. + prompt_cache.clear() + output_cache.clear() + _validate_history_sources(history) + if ( + _source_signature(source) != signature + or not source_matches_context(source) + or ( + request_context is not None + and _render_context_key( + list( + _request_messages( + cast( + ChatCompletionsExchange + | MessagesExchange + | ResponsesExchange, + exchange, + ) + ) + ) + ) + != request_context + ) + ): + raise ValueError( + "Sampled source changed while proving recorded request roles" + ) break canonical_length_stop_mask = _synthetic_length_stop_mask( @@ -5407,32 +5612,48 @@ def source_matches_context(source: object) -> bool: canonical_assistant_mask, direct_bounds or None, ) - # These four masks translate the same pair of token sequences. - mask_opcodes: list[tuple[str, int, int, int, int]] = [] - assistant_mask = _translate_token_mask( - canonical_rendered, - rendered, - canonical_assistant_mask, - tokenizer=resolved_tokenizer, - _opcodes=mask_opcodes, - ) - output_mask = _translate_token_mask( - canonical_rendered, - rendered, - canonical_output_mask, - tokenizer=resolved_tokenizer, - _opcodes=mask_opcodes, - ) - stop_mask = _translate_token_mask( - canonical_rendered, rendered, canonical_stop_mask, _opcodes=mask_opcodes - ) - length_stop_mask = _translate_token_mask( - canonical_rendered, - rendered, - canonical_length_stop_mask, - _opcodes=mask_opcodes, - ) - mask_opcodes.clear() + if recorded_prompt_masks is not None and not any( + canonical_output_mask[:canonical_prefix_length] + + canonical_length_stop_mask[:canonical_prefix_length] + ): + assistant_prefix, stop_prefix = recorded_prompt_masks + assistant_mask = ( + assistant_prefix + canonical_assistant_mask[canonical_prefix_length:] + ) + stop_mask = stop_prefix + canonical_stop_mask[canonical_prefix_length:] + output_mask = [False] * exact_prefix_length + canonical_output_mask[ + canonical_prefix_length: + ] + length_stop_mask = [False] * exact_prefix_length + canonical_length_stop_mask[ + canonical_prefix_length: + ] + else: + # These four masks translate the same pair of token sequences. + mask_opcodes: list[tuple[str, int, int, int, int]] = [] + assistant_mask = _translate_token_mask( + canonical_rendered, + rendered, + canonical_assistant_mask, + tokenizer=resolved_tokenizer, + _opcodes=mask_opcodes, + ) + output_mask = _translate_token_mask( + canonical_rendered, + rendered, + canonical_output_mask, + tokenizer=resolved_tokenizer, + _opcodes=mask_opcodes, + ) + stop_mask = _translate_token_mask( + canonical_rendered, rendered, canonical_stop_mask, _opcodes=mask_opcodes + ) + length_stop_mask = _translate_token_mask( + canonical_rendered, + rendered, + canonical_length_stop_mask, + _opcodes=mask_opcodes, + ) + mask_opcodes.clear() positions_by_first_token: dict[int, list[int]] = {} for index, token_id in enumerate(rendered): positions_by_first_token.setdefault(token_id, []).append(index) @@ -6034,6 +6255,23 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: ) ) ): + # The exact builder owns sampled spans, but the rendered request + # already proved role labels before the first sampled response. + # Retain those labels instead of discarding them on this return. + if ( + exact_prefix_length + and exact.tokens[:exact_prefix_length] == rendered[:exact_prefix_length] + and not any( + flag & TokenFlag.SAMPLED + for flag in exact.flags[:exact_prefix_length] + ) + ): + for index in range(exact_prefix_length): + exact.flags[index] |= _rendered_flag( + assistant_mask[index] and not length_stop_mask[index], + output_mask[index] and not length_stop_mask[index], + stop_mask[index], + ) return exact sampled_message_count = sum( diff --git a/tests/unit/trajectories/test_evidence_reuse.py b/tests/unit/trajectories/test_evidence_reuse.py new file mode 100644 index 000000000..72e749b56 --- /dev/null +++ b/tests/unit/trajectories/test_evidence_reuse.py @@ -0,0 +1,257 @@ +from __future__ import annotations + +from collections import Counter +import copy +from typing import Any, cast + +import pytest +from test_tokenize import _chat_exchange + +import art.trajectories as tr +from art.trajectories import _tokenize as module + + +def trajectory(*exchanges): + return tr.Trajectory( + exchanges=tr.TrajectoryExchanges(chat_completions=list(exchanges)) + ) + + +def logprobs(exchange): + value = exchange.response.choices[0].logprobs + assert value is not None and value.content is not None + return value.content + + +def extras(exchange): + value = exchange.response.choices[0].model_extra + assert value is not None + return value + + +def sources(history): + return [s for s in history.message_sources if module._source_is_sampled(s)] + + +def test_exact_assembly_reuses_evidence_only_within_one_call(monkeypatch): + first = _chat_exchange([1], [2, 3]) + second = _chat_exchange([1, 2, 3, 4], [5, 6], offset=1) + first.response.choices[0].index = 7 + value = trajectory(first, second) + calls = [] + fingerprint = module._fingerprint + + def observed(evidence): + calls.append(evidence) + return fingerprint(evidence) + + monkeypatch.setattr(module, "_fingerprint", observed) + before = value.model_dump_json() + result = value.tokenize() + assert result.tokens == [1, 2, 3, 4, 5, 6] + assert len(calls) == 4 # Two preflight fingerprints and two assembly fingerprints. + assert value.model_dump_json() == before + lp = logprobs(first) + lp[1].logprob = -7.5 + assert value.tokenize().logprobs[2] == -7.5 + assert result.logprobs[2] == -0.3 + assert len(calls) == 8 + + +@pytest.mark.parametrize( + "field", ["prompt", "output", "logprob", "content", "reason", "index"] +) +def test_signature_remains_fresh_after_source_edit(field): + exchange = _chat_exchange([1], [2, 3]) + source = sources(trajectory(exchange).histories()[0])[0] + before = module._source_signature(source) + choice = exchange.response.choices[0] + if field == "prompt": + extras(exchange)["prompt_token_ids"][0] = 9 + elif field == "output": + extras(exchange)["token_ids"][0] = 9 + elif field == "logprob": + logprobs(exchange)[0].logprob = float("nan") + elif field == "content": + choice.message.content = "edited public view" + elif field == "reason": + choice.finish_reason = "length" + else: + choice.index = 7 + source = source.model_copy(update={"choice_index": 7}) + assert module._source_signature(source) != before + + +def test_source_match_caches_exchange_identity_not_response_id(monkeypatch): + first = _chat_exchange([1], [2, 3]) + source = sources(trajectory(first).histories()[0])[0] + same = copy.copy(source) + different = copy.deepcopy(source) + calls = [] + fingerprint = module._fingerprint + + def observed(evidence): + calls.append(evidence) + return fingerprint(evidence) + + monkeypatch.setattr(module, "_fingerprint", observed) + assert module._sources_match([source, same], [same, source]) + assert len(calls) == 1 + assert module._sources_match([source], [different]) + assert len(calls) == 3 + different.exchange.response.choices[0].logprobs.content[0].logprob = -9 + assert not module._sources_match([source], [different]) + assert len(calls) == 5 + source.exchange.response.choices[0].message.content = "fresh mutation" + assert not module._sources_match([source], [different]) + assert len(calls) == 7 + + +def test_fingerprint_cache_bound_and_identity(): + cache: dict[tuple[int, str, int], tuple[module.Exchange, str]] = {} + first = _chat_exchange([1], [2]) + expected = module._sampled_evidence_fingerprint( + first, protocol="chat_completions", index=0 + ) + for i in range(258): + exchange = _chat_exchange([i], [i + 1]) + module._sampled_evidence_fingerprint( + exchange, protocol="chat_completions", index=0, _cache=cache + ) + assert len(cache) == 256 + assert ( + module._sampled_evidence_fingerprint( + first, protocol="chat_completions", index=0, _cache=cache + ) + == expected + ) + # Even a stale identity slot cannot borrow another exchange's evidence. + cache = {(id(first), "chat_completions", 0): (exchange, "wrong")} + assert ( + module._sampled_evidence_fingerprint( + first, protocol="chat_completions", index=0, _cache=cache + ) + == expected + ) + + +def test_whitespace_decoder_clears_evidence_and_record_caches(): + first = _chat_exchange([1], [2, 3]) + first.response.choices[0].finish_reason = "length" + second = _chat_exchange([1, 2, 3, 32, 9, 8], [4, 5], offset=1) + third = _chat_exchange([1, 2, 3, 32, 9, 8, 4, 5, 7], [6], offset=2) + history = trajectory(first, second, third).histories()[0] + first_source, second_source, _ = sources(history) + second_key = module._sampled_source_key(second_source) + nested = trajectory(_chat_exchange([88], [99])) + + class Decoder: + eos_token_id = 9 + all_special_ids = [] + calls = 0 + + def decode(self, ids): + assert ids == [32] + self.calls += 1 + assert nested.tokenize().tokens == [88, 99] + logprobs(second)[0].logprob = -7.5 + return " " + + decoder = Decoder() + trace = module._TraceBuilder() + value = module._tokenize_exact_projected_chat_history( + history, + tokenizer=cast(Any, decoder), + projection_validated=True, + _trace=trace, + length_stop_boundaries={ + module._sampled_source_key( + first_source + ): module._RenderedLengthStopBoundary(tail=(9,), following=(8,)) + }, + ) + assert value is not None and decoder.calls == 1 + assert trace.trace is not None + assert value.logprobs[6] == -7.5 + assert trace.trace.source_keys[6] == module._sampled_source_key(second_source) + assert trace.trace.source_keys[6] != second_key + + +def test_copied_context_stop_callback_clears_evidence_and_records(): + first = _chat_exchange([1], [2, 3]) + extras(first)["stop_reason"] = "public-stop" + second = _chat_exchange([1, 3, 4], [5, 6], offset=1) + third = _chat_exchange([1, 3, 4, 5, 6, 7], [8], offset=2) + histories = trajectory(first, second, third).histories() + assert len(histories) == 2 + original_trace = module._TraceBuilder() + original = module._tokenize_exact_projected_chat_history( + histories[0], tokenizer=None, projection_validated=True, _trace=original_trace + ) + assert original is not None and original_trace.trace is not None + second_source = next(s for s in sources(histories[1]) if s.exchange is second) + old_key = module._sampled_source_key(second_source) + + class StopTokenizer: + eos_token_id = 3 + all_special_ids = [] + calls = 0 + + def __call__(self, text, **kwargs): + assert text == "public-stop" + self.calls += 1 + logprobs(second)[0].logprob = -8.5 + return {"input_ids": [3]} + + decoder = StopTokenizer() + trace = module._TraceBuilder() + value = module._tokenize_exact_projected_chat_history( + histories[1], + tokenizer=cast(Any, decoder), + projection_validated=True, + _trace=trace, + _strict_sources=True, + _prior=[(original, original_trace.trace)], + ) + assert value is not None and decoder.calls == 1 + assert trace.trace is not None + assert value.logprobs[3] == -8.5 + assert trace.trace.source_keys[3] == module._sampled_source_key(second_source) + assert trace.trace.source_keys[3] != old_key + assert not value.flags[1] & tr.TokenFlag.SAMPLED + + +def test_decoder_exception_identity_and_next_call_fresh(): + first = _chat_exchange([1], [2, 3]) + extras(first)["stop_reason"] = "public-stop" + value = trajectory(first) + failure = ValueError("public stop decoder failed") + + class Broken: + def __call__(self, text, **kwargs): + raise failure + + with pytest.raises(ValueError) as caught: + value.tokenize(tokenizer=Broken()) + assert caught.value is failure + logprobs(first)[0].logprob = -10 + assert value.tokenize().logprobs[1] == -10 + + +@pytest.mark.parametrize("edit", ["model", "source", "prompt"]) +def test_edited_history_does_not_reuse_projection_proof(edit): + from test_tokenize import _CharacterTemplateTokenizer + + value = trajectory(_chat_exchange([1], [2, 3])) + history = value.histories()[0] + history.tokenize() + if edit == "model": + history.model = "different/model" + elif edit == "source": + history.message_sources[-1] = history.message_sources[-1].model_copy( + update={"choice_index": 99} + ) + else: + history.messages[0]["content"] = "new question" + with pytest.raises((ValueError, AssertionError)): + history.tokenize(tokenizer=_CharacterTemplateTokenizer()) diff --git a/tests/unit/trajectories/test_recorded_prompt_roles.py b/tests/unit/trajectories/test_recorded_prompt_roles.py new file mode 100644 index 000000000..85bdcd387 --- /dev/null +++ b/tests/unit/trajectories/test_recorded_prompt_roles.py @@ -0,0 +1,432 @@ +from __future__ import annotations + +from copy import deepcopy +import json +import math +from typing import Any, cast + +from openai.types.chat import ChatCompletionMessageParam +import pytest +from test_literal_thinking_off import _TEMPLATE, _history + +import art.trajectories as tr +from art.trajectories import _tokenize as module + + +def _extras(exchange: tr.ChatCompletionsExchange) -> dict[str, Any]: + extra = exchange.response.choices[0].model_extra + assert extra is not None + return extra + + +def _case(content: str, *, output: str = "New recorded answer"): + history, tokenizer = _history(content=output) + source = history.message_sources[-1] + assert source is not None + exchange = source.exchange + assert isinstance(exchange, tr.ChatCompletionsExchange) + messages = [ + {"role": "user", "content": "Old public query"}, + {"role": "assistant", "content": content}, + {"role": "user", "content": "New public query"}, + ] + exchange.request["messages"] = cast( + list[ChatCompletionMessageParam], deepcopy(messages) + ) + prompt = tokenizer.apply_chat_template( + messages, + chat_template=_TEMPLATE, + add_generation_prompt=True, + enable_thinking=False, + preserve_thinking=True, + ) + _extras(exchange)["prompt_token_ids"] = prompt + trajectory = tr.Trajectory( + exchanges=tr.TrajectoryExchanges(chat_completions=[exchange]) + ) + return trajectory, tokenizer, prompt, exchange + + +@pytest.mark.parametrize( + "content", + [ + "Public reasoning\n\n\n\n\nPublic answer", + "Public reasoningPublic answer", + "Public reasoningmiddlePublic answer", + ], +) +def test_recorded_request_roles_follow_proved_historical_renderer( + monkeypatch: pytest.MonkeyPatch, content: str +) -> None: + trajectory, tokenizer, prompt, exchange = _case(content) + before = trajectory.model_dump() + kwargs = dict(tokenizer=tokenizer, multi_history=True) + result = trajectory.tokenize(**kwargs) + assert len(result.histories) == 1 + actual = result.histories[0] + output = _extras(exchange)["token_ids"] + assert actual.tokens[: len(prompt)] == prompt + assert actual.tokens[len(prompt) : len(prompt) + len(output)] == output + assert actual.logprobs[len(prompt) : len(prompt) + len(output)] == [-0.5] * len( + output + ) + assert all(math.isnan(lp) for lp in actual.logprobs[: len(prompt)]) + required = ( + tr.TokenFlag.SAMPLED + | tr.TokenFlag.EXACT + | tr.TokenFlag.ASSISTANT + | tr.TokenFlag.OUTPUT + ) + assert all( + flag & required == required + for flag in actual.flags[len(prompt) : len(prompt) + len(output)] + ) + assert not any( + flag & (tr.TokenFlag.SAMPLED | tr.TokenFlag.OUTPUT) + for flag in actual.flags[: len(prompt)] + ) + assert any(flag & tr.TokenFlag.ASSISTANT for flag in actual.flags[: len(prompt)]) + assert sum( + tr.first_occurrence_masks(result.histories, where=tr.TokenFlag.SAMPLED)[0] + ) == len(output) + assert trajectory.model_dump() == before + # This original template is the successful historical rendering oracle for + # this public request; new source normalization must not change its roles. + with monkeypatch.context() as patch: + patch.setattr( + module, "chat_template_with_preserved_thinking", lambda value: value + ) + historical = trajectory.tokenize(**kwargs).histories[0] + assert actual.tokens == historical.tokens + assert actual.flags == historical.flags + assert all( + a == b or math.isnan(a) and math.isnan(b) + for a, b in zip(actual.logprobs, historical.logprobs, strict=True) + ) + for flag in ( + tr.TokenFlag.SAMPLED, + tr.TokenFlag.OUTPUT, + tr.TokenFlag.ASSISTANT, + tr.TokenFlag.STOP, + ): + assert tr.first_occurrence_masks( + [actual], where=flag + ) == tr.first_occurrence_masks([historical], where=flag) + + +def test_historical_context_does_not_reenable_literal_parsing_for_new_output() -> None: + literal = "New literal must remain content" + trajectory, tokenizer, prompt, exchange = _case( + "Public reasoning\n\n\n\n\nPublic answer", output=literal + ) + result = trajectory.tokenize(tokenizer=tokenizer, multi_history=True) + actual = result.histories[0] + output = _extras(exchange)["token_ids"] + assert actual.tokens[len(prompt) : len(prompt) + len(output)] == output + assert literal in tokenizer.rendered[-1] or any( + literal in text for text in tokenizer.rendered + ) + + +def test_historical_tool_serialization_uses_original_request_key_order( + monkeypatch: pytest.MonkeyPatch, +) -> None: + trajectory, tokenizer, _, exchange = _case( + "Public reasoning\n\n\n\n\nPublic answer" + ) + # HF's template tojson filter preserves insertion order. + tokenizer.env.policies["json.dumps_kwargs"] = {"sort_keys": False} + tools = [ + { + "type": "function", + "function": { + "parameters": {"type": "object", "properties": {}}, + "description": "Public tool", + "name": "lookup", + }, + } + ] + exchange.request["tools"] = tools + prompt = tokenizer.apply_chat_template( + exchange.request["messages"], + tools=tools, + chat_template=_TEMPLATE, + add_generation_prompt=True, + enable_thinking=False, + preserve_thinking=True, + ) + _extras(exchange)["prompt_token_ids"] = prompt + history = trajectory.chat_completions_history() + assert history.tools == tools + assert json.dumps(history.tools) != json.dumps(tools) + assert ( + tokenizer.apply_chat_template( + history.messages[:-1], + tools=history.tools, + chat_template=_TEMPLATE, + add_generation_prompt=True, + enable_thinking=False, + preserve_thinking=True, + ) + != prompt + ) + before = trajectory.model_dump() + actual = trajectory.tokenize(tokenizer=tokenizer, multi_history=True).histories[0] + with monkeypatch.context() as patch: + patch.setattr( + module, "chat_template_with_preserved_thinking", lambda value: value + ) + historical = trajectory.tokenize( + tokenizer=tokenizer, multi_history=True + ).histories[0] + assert actual.tokens == historical.tokens + assert actual.flags == historical.flags + assert actual.tokens[: len(prompt)] == prompt + assert trajectory.model_dump() == before + + +@pytest.mark.parametrize("change", ["prompt", "edited", "override"]) +def test_unproved_historical_renderer_does_not_certify_context( + monkeypatch: pytest.MonkeyPatch, change: str +) -> None: + trajectory, tokenizer, _, exchange = _case( + "Public reasoning\n\n\n\n\nPublic answer" + ) + history = trajectory.chat_completions_history() + kwargs: dict[str, Any] = {"tokenizer": tokenizer} + if change == "prompt": + _extras(exchange)["prompt_token_ids"][0] += 1 + elif change == "edited": + history.messages[1]["content"] = "Edited public context" + else: + kwargs["chat_template"] = _TEMPLATE + "{# explicit renderer #}" + original = module._recorded_prompt_role_masks + admitted = [] + + def observe(*args: Any, **options: Any): + value = original(*args, **options) + admitted.append(value) + return value + + monkeypatch.setattr(module, "_recorded_prompt_role_masks", observe) + + def outcome(): + try: + return history.tokenize(**kwargs) + except ValueError as error: + return type(error), str(error) + + actual = outcome() + assert not any(value is not None for value in admitted) + monkeypatch.setattr( + module, "_recorded_prompt_role_masks", lambda *args, **kwargs: None + ) + baseline = outcome() + if isinstance(actual, tuple): + assert actual == baseline + else: + assert actual.tokens == baseline.tokens + assert actual.flags == baseline.flags + + +def test_sampled_history_is_not_reclassified_as_request_context() -> None: + history, tokenizer = _history() + assert ( + module._recorded_prompt_role_masks( + cast(list[dict[str, Any]], history.messages), + history.message_sources, + [], + tokenizer=tokenizer, + template=_TEMPLATE, + tools=None, + kwargs={}, + ) + is None + ) + + +def test_unchanged_template_length_retry_preserves_request_roles( + monkeypatch: pytest.MonkeyPatch, +) -> None: + trajectory, tokenizer, prompt, _ = _case("Plain historical assistant content") + monkeypatch.setattr( + module, "chat_template_with_preserved_thinking", lambda value: value + ) + result = trajectory.tokenize(tokenizer=tokenizer, multi_history=True).histories[0] + assert result.tokens[: len(prompt)] == prompt + assert any(flag & tr.TokenFlag.ASSISTANT for flag in result.flags[: len(prompt)]) + assert not any( + flag & (tr.TokenFlag.OUTPUT | tr.TokenFlag.SAMPLED) + for flag in result.flags[: len(prompt)] + ) + + +@pytest.mark.parametrize( + "change", + ["message", "tools", "kwargs", "source", "tool_order", "source_tool_order"], +) +def test_historical_proof_refuses_renderer_mutation( + monkeypatch: pytest.MonkeyPatch, change: str +) -> None: + trajectory, tokenizer, _, exchange = _case( + "Public reasoning\n\n\n\n\nPublic answer" + ) + render = tokenizer.apply_chat_template + calls = 0 + + def mutate(messages, **kwargs): + nonlocal calls + value = render(messages, **kwargs) + if kwargs.get("chat_template") == _TEMPLATE: + calls += 1 + if calls == 2: + if change == "message": + messages[0]["content"] = "Changed public content" + elif change == "source": + _extras(exchange)["prompt_token_ids"][0] += 1 + elif change == "tools": + kwargs["tools"][0]["function"]["description"] = "Changed" + elif change in {"tool_order", "source_tool_order"}: + tools = ( + kwargs["tools"] + if change == "tool_order" + else exchange.request["tools"] + ) + function = tools[0]["function"] + function["name"] = function.pop("name") + else: + # Mutate nested kwargs, which Python's ** expansion shares. + kwargs["public_option"]["changed"] = True + return value + + if change in {"tools", "tool_order", "source_tool_order"}: + exchange.request["tools"] = [ + { + "type": "function", + "function": { + "name": "lookup", + "description": "Original", + "parameters": {}, + }, + } + ] + # Tools affect this template, so reconstruct the authoritative prompt. + _extras(exchange)["prompt_token_ids"] = render( + exchange.request["messages"], + tools=exchange.request["tools"], + chat_template=_TEMPLATE, + add_generation_prompt=True, + enable_thinking=False, + preserve_thinking=True, + ) + if change == "kwargs": + exchange.request["chat_template_kwargs"]["public_option"] = {"changed": False} + monkeypatch.setattr(tokenizer, "apply_chat_template", mutate) + with pytest.raises(ValueError, match="changed|does not match|differs"): + trajectory.tokenize(tokenizer=tokenizer, multi_history=True) + assert calls >= 2 + + +@pytest.mark.parametrize( + "error", [RuntimeError("public callback"), KeyboardInterrupt(), SystemExit(9)] +) +def test_historical_proof_preserves_callback_exception_identity( + monkeypatch: pytest.MonkeyPatch, error: BaseException +) -> None: + trajectory, tokenizer, _, _ = _case( + "Public reasoning\n\n\n\n\nPublic answer" + ) + render = tokenizer.apply_chat_template + + def fail(messages, **kwargs): + if kwargs.get("chat_template") == _TEMPLATE: + raise error + return render(messages, **kwargs) + + monkeypatch.setattr(tokenizer, "apply_chat_template", fail) + with pytest.raises(type(error)) as caught: + trajectory.tokenize(tokenizer=tokenizer, multi_history=True) + assert caught.value is error + + +@pytest.mark.parametrize("error", [TypeError(), KeyError(), NotImplementedError()]) +def test_unsupported_historical_prefix_keeps_existing_translation( + monkeypatch: pytest.MonkeyPatch, error: Exception +) -> None: + trajectory, tokenizer, _, _ = _case("Public reasoningPublic answer") + with monkeypatch.context() as patch: + patch.setattr( + module, "_recorded_prompt_role_masks", lambda *args, **kwargs: None + ) + expected = trajectory.tokenize(tokenizer=tokenizer, multi_history=True) + render = tokenizer.apply_chat_template + failures = 0 + + def unsupported(messages, **kwargs): + nonlocal failures + if kwargs.get("chat_template") == _TEMPLATE and len(messages) < 3: + failures += 1 + raise error + return render(messages, **kwargs) + + monkeypatch.setattr(tokenizer, "apply_chat_template", unsupported) + actual = trajectory.tokenize(tokenizer=tokenizer, multi_history=True) + assert failures == 1 + assert actual.histories[0].tokens == expected.histories[0].tokens + assert actual.histories[0].flags == expected.histories[0].flags + + +@pytest.mark.parametrize("error", [TypeError(), NotImplementedError()]) +def test_unsupported_historical_offsets_keep_existing_translation( + monkeypatch: pytest.MonkeyPatch, error: Exception +) -> None: + trajectory, tokenizer, _, _ = _case("Public reasoningPublic answer") + with monkeypatch.context() as patch: + patch.setattr( + module, "_recorded_prompt_role_masks", lambda *args, **kwargs: None + ) + expected = trajectory.tokenize(tokenizer=tokenizer, multi_history=True) + encode = type(tokenizer).__call__ + failures = 0 + + def unsupported(self, text, **kwargs): + nonlocal failures + if kwargs.get("return_offsets_mapping"): + failures += 1 + raise error + return encode(self, text, **kwargs) + + monkeypatch.setattr(type(tokenizer), "__call__", unsupported) + actual = trajectory.tokenize(tokenizer=tokenizer, multi_history=True) + assert failures >= 1 + assert actual.histories[0].tokens == expected.histories[0].tokens + assert actual.histories[0].flags == expected.histories[0].flags + + +@pytest.mark.parametrize("offset", [(0, 0), (-1, 1), (0, 100000), (False, 1)]) +def test_unproved_historical_offset_declines( + monkeypatch: pytest.MonkeyPatch, offset: tuple[int, int] +) -> None: + trajectory, tokenizer, prompt, _ = _case("Public reasoningPublic answer") + history = trajectory.chat_completions_history() + encode = type(tokenizer).__call__ + + def malformed(self, text, **kwargs): + result = encode(self, text, **kwargs) + if kwargs.get("return_offsets_mapping"): + result["offset_mapping"][0] = offset + return result + + monkeypatch.setattr(type(tokenizer), "__call__", malformed) + assert ( + module._recorded_prompt_role_masks( + history.messages[:-1], + history.message_sources[:-1], + prompt, + tokenizer=tokenizer, + template=_TEMPLATE, + tools=None, + kwargs={"enable_thinking": False, "preserve_thinking": True}, + ) + is None + ) From 9b1507b6c3792dd2416b576ff75efc2b125194f3 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 04:46:46 +0000 Subject: [PATCH 11/16] Reuse native source evidence across the tokenization stop decision --- src/art/trajectories/_tokenize.py | 26 +++++- .../unit/trajectories/test_evidence_reuse.py | 88 ++++++++++++++++++- 2 files changed, 108 insertions(+), 6 deletions(-) diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index 1ae9cdc8f..42d3a68e4 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -3393,7 +3393,11 @@ def _last_source_exchange(sources: Sequence[object]) -> Exchange | None: return None -def _history_has_length_stop(history: History) -> bool: +def _history_has_length_stop( + history: History, + *, + _fingerprints: dict[tuple[int, str, int], tuple[Exchange, str]] | None = None, +) -> bool: sources: Sequence[object] if isinstance(history, (ChatCompletionsHistory, AnthropicMessagesHistory)): sources = history.message_sources @@ -3405,7 +3409,7 @@ def _history_has_length_stop(history: History) -> bool: for source in sources: if source is None or not _source_is_sampled(source): continue - source_key = _sampled_source_key(source) + source_key = _sampled_source_key(source, _fingerprints=_fingerprints) if source_key in seen: continue seen.add(source_key) @@ -4533,11 +4537,14 @@ def _tokenize_exact_projected_chat_history( _trace: _TraceBuilder | None = None, _strict_sources: bool = False, _prior: Sequence[tuple[TokenizedHistory, _HistoryTokenizationTrace]] = (), + _fingerprints: dict[tuple[int, str, int], tuple[Exchange, str]] | None = None, ) -> TokenizedHistory | None: if not projection_validated and not _history_matches_projection(history): return None # Reuse evidence only in this callback-free phase, never across render/decode. - fingerprints: dict[tuple[int, str, int], tuple[Exchange, str]] = {} + fingerprints: dict[tuple[int, str, int], tuple[Exchange, str]] = ( + {} if _fingerprints is None else _fingerprints + ) sampled_sources: list[object] = [] seen: set[tuple[object, ...]] = set() for message, source in zip(history.messages, history.message_sources, strict=True): @@ -7388,7 +7395,17 @@ def _tokenize_history( can_render = tokenizer is None or callable( getattr(tokenizer, "apply_chat_template", None) ) - has_length_stop = can_render and _history_has_length_stop(history) + # Without a tokenizer, the stop decision and first exact assembly have no + # user callback between them. Keep their evidence in one bounded phase; + # never carry it into a rendered or tokenizer-supplied path. + fingerprints: dict[tuple[int, str, int], tuple[Exchange, str]] | None = ( + {} + if tokenizer is None and isinstance(history, ChatCompletionsHistory) + else None + ) + has_length_stop = can_render and _history_has_length_stop( + history, _fingerprints=fingerprints + ) needs_synthetic_stop = _history_needs_synthetic_stop(history, tokenizer) needs_render = ( render_state.needs_render @@ -7422,6 +7439,7 @@ def _tokenize_history( _trace=_trace, _strict_sources=True, _prior=_prior, + _fingerprints=fingerprints, ) ) ): diff --git a/tests/unit/trajectories/test_evidence_reuse.py b/tests/unit/trajectories/test_evidence_reuse.py index 72e749b56..a1dc86d8a 100644 --- a/tests/unit/trajectories/test_evidence_reuse.py +++ b/tests/unit/trajectories/test_evidence_reuse.py @@ -49,13 +49,97 @@ def observed(evidence): before = value.model_dump_json() result = value.tokenize() assert result.tokens == [1, 2, 3, 4, 5, 6] - assert len(calls) == 4 # Two preflight fingerprints and two assembly fingerprints. + assert len(calls) == 2 # One callback-free decision/assembly phase per source. assert value.model_dump_json() == before lp = logprobs(first) lp[1].logprob = -7.5 assert value.tokenize().logprobs[2] == -7.5 assert result.logprobs[2] == -0.3 - assert len(calls) == 8 + assert len(calls) == 4 + + +def test_supplied_tokenizer_stop_probe_cannot_lend_stale_evidence(monkeypatch): + first = _chat_exchange([1], [2, 3]) + extras(first)["stop_reason"] = "public-stop" + second = _chat_exchange([1, 2, 3, 4], [5, 9], offset=1) + value = trajectory(first, second) + history = value.histories()[0] + second_source = sources(history)[1] + original_key = module._sampled_source_key(second_source) + + class Tokenizer: + eos_token_id = 9 + all_special_ids = [] + calls = 0 + + def apply_chat_template(self, *args, **kwargs): + raise AssertionError("complete records need no rendering") + + def __call__(self, text, **kwargs): + assert text == "public-stop" + self.calls += 1 + logprobs(second)[0].logprob = -8.5 + return {"input_ids": [3]} + + tokenizer = Tokenizer() + trace = module._TraceBuilder() + result = module.tokenize_history( + history, + model=history.model, + base_model=None, + tokenizer=cast(Any, tokenizer), + chat_template=None, + chat_template_kwargs=None, + _trace=trace, + ) + assert tokenizer.calls >= 1 and trace.trace is not None + assert result.tokens == [1, 2, 3, 4, 5, 9] + assert result.logprobs[4] == -8.5 + assert trace.trace.source_keys[4] == module._sampled_source_key(second_source) + assert trace.trace.source_keys[4] != original_key + assert result.flags[2] & tr.TokenFlag.STOP + assert result.flags[-1] & tr.TokenFlag.STOP + + +@pytest.mark.parametrize("override", [False, True]) +def test_render_fallback_does_not_receive_decision_evidence(monkeypatch, override): + from test_tokenize import _character_template_history + + history, tokenizer, _ = _character_template_history() + first_source = sources(history)[0] + exchange = first_source.exchange + old_key = module._sampled_source_key(first_source) + inner = trajectory(_chat_exchange([88], [99])) + + def load(config): + # A nested tokenization and a source edit happen after the original + # length decision, at an existing renderer-loader callback boundary. + assert inner.tokenize().tokens == [88, 99] + logprobs(exchange)[0].logprob = -7.5 + return tokenizer + + monkeypatch.setattr(module, "_load_tokenizer", load) + monkeypatch.setattr( + module, + "_tokenizer_config", + lambda *args: module._TokenizerConfig("public/base"), + ) + trace = module._TraceBuilder() + result = module.tokenize_history( + history, + model=history.model, + base_model="public/base", + tokenizer=None, + chat_template="explicit public template" if override else None, + chat_template_kwargs=None, + _trace=trace, + ) + assert trace.trace is not None + new_key = module._sampled_source_key(first_source) + assert new_key != old_key + indices = [i for i, key in enumerate(trace.trace.source_keys) if key == new_key] + assert indices and result.logprobs[indices[0]] == -7.5 + assert old_key not in trace.trace.source_keys @pytest.mark.parametrize( From f62857b81a58a36597997cb4620b3578965121b5 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 05:10:39 +0000 Subject: [PATCH 12/16] Isolate the warning state in the callback fallback regression --- tests/unit/trajectories/test_evidence_reuse.py | 1 + 1 file changed, 1 insertion(+) diff --git a/tests/unit/trajectories/test_evidence_reuse.py b/tests/unit/trajectories/test_evidence_reuse.py index a1dc86d8a..8555ffe96 100644 --- a/tests/unit/trajectories/test_evidence_reuse.py +++ b/tests/unit/trajectories/test_evidence_reuse.py @@ -105,6 +105,7 @@ def __call__(self, text, **kwargs): def test_render_fallback_does_not_receive_decision_evidence(monkeypatch, override): from test_tokenize import _character_template_history + monkeypatch.setattr(module, "_WARNED_PREFIX_RETOKENIZATION", False) history, tokenizer, _ = _character_template_history() first_source = sources(history)[0] exchange = first_source.exchange From b310d69f71c3e3144f12f8cebf8035016630b213 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 06:33:01 +0000 Subject: [PATCH 13/16] Use final recorded output as the end of complete native Chat histories --- docs/features/additional-histories.mdx | 10 +- src/art/trajectories/_tokenize.py | 26 +- .../unit/trajectories/test_native_terminal.py | 236 ++++++++++++++++++ .../trajectories/test_recorded_boundaries.py | 19 +- tests/unit/trajectories/test_tokenize.py | 33 +-- 5 files changed, 292 insertions(+), 32 deletions(-) create mode 100644 tests/unit/trajectories/test_native_terminal.py diff --git a/docs/features/additional-histories.mdx b/docs/features/additional-histories.mdx index 433f6ec68..154c9fd90 100644 --- a/docs/features/additional-histories.mdx +++ b/docs/features/additional-histories.mdx @@ -161,7 +161,15 @@ existing protocol-specific projection rules; no separate tokenization API or native-representation option is needed. `multi_history=True` preserves the histories selected by the trajectory, including their order and model selection. -Templates still own unrecorded separators, role masks, and synthetic stop tokens. +Complete unchanged Chat histories end at the final recorded response token, +including tool responses and responses stopped by a length limit. ART does not +reconstruct that sampled body from its text or structured tool projection, or add +an unobserved terminal footer. This deliberately excludes synthetic terminal +tokens from `OUTPUT`/SFT masks; recorded sampled tokens and logprobs are unchanged. +Explicit template overrides and incomplete or edited histories retain rendering. + +Templates still prove nonterminal separators, role masks, and synthetic stop +tokens against the next recorded prompt. If tokenizing another history in the same trajectory already resolves a tokenizer for the same model, ART reuses that authority to label recorded sampled stop tokens. This does not trigger a new tokenizer load or change rendering. Complete diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index 42d3a68e4..61626c6fc 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -4869,7 +4869,9 @@ def _tokenize_recorded_chat_boundaries( terminators = _terminator_ids(tokenizer) if not terminators: return None - for ordinal, (index, source, prompt, output, _) in enumerate(entries): + # The last native output ends the recorded history. No later prompt proves + # an additional footer, so do not reconstruct one from a lossy projection. + for ordinal, (index, source, prompt, output, _) in enumerate(entries[:-1]): key = _sampled_source_key(source) stop, _ = _source_stop_evidence(source, key) if stop not in {"stop", "length"}: @@ -6144,6 +6146,9 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: assert source is not None source_key = _sampled_source_key(source) stop_reason = _source_stop_evidence(source, source_key)[0] + if _recorded_boundaries and position + 1 == len(sampled_message_indices): + length_stop_count += stop_reason == "length" + continue output = _source_output_tokens(source, source_key) synthetic_stop = ( stop_reason == "stop" @@ -6250,7 +6255,7 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: continue length_stop_boundaries[source_key] = boundary if ( - length_stop_count + (length_stop_count or _recorded_boundaries) and length_stop_boundaries_complete and ( exact := _tokenize_exact_projected_chat_history( @@ -7424,7 +7429,22 @@ def _tokenize_history( return exact if isinstance(history, ChatCompletionsHistory): if ( - not has_length_stop + ( + not has_length_stop + or not _copied_context + and sum( + message.get("role") == "assistant" for message in history.messages + ) + == 1 + and all( + message.get("role") != "assistant" + or source is not None + and _source_is_sampled(source) + for message, source in zip( + history.messages, history.message_sources, strict=True + ) + ) + ) and not needs_synthetic_stop and not override_requires_render and not render_state.context_changed diff --git a/tests/unit/trajectories/test_native_terminal.py b/tests/unit/trajectories/test_native_terminal.py new file mode 100644 index 000000000..7858806aa --- /dev/null +++ b/tests/unit/trajectories/test_native_terminal.py @@ -0,0 +1,236 @@ +from __future__ import annotations + +from copy import deepcopy +import json +import math +from typing import Any, cast + +from openai.types.chat import ChatCompletion, ChatCompletionMessageParam +import pytest +from test_tokenize import ( + _character_template_history, + _CharacterTemplateTokenizer, + _chat_exchange, +) + +import art.trajectories as tr +from art.trajectories import _tokenize as module + + +@pytest.fixture(autouse=True) +def restore_warning_state(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(module, "_WARNED_PREFIX_RETOKENIZATION", False) + + +class ProjectedToolTokenizer(_CharacterTemplateTokenizer): + def apply_chat_template( + self, + messages: Any, + *, + tokenize: bool = True, + add_generation_prompt: bool, + **kwargs: Any, + ) -> Any: + text = "" + for message in messages: + text += message["role"] + ":" + str(message.get("content") or "") + if message.get("tool_calls"): + text += json.dumps(message["tool_calls"], sort_keys=True) + if message["role"] == "assistant": + text += "§" + if add_generation_prompt: + text += "assistant:" + return self._encode(text) if tokenize else text + + +def projected_history( + *, finish: str, sampled_eos: bool, earlier_length: bool +) -> tuple[ + tr.Trajectory, + ProjectedToolTokenizer, + list[tuple[list[int], list[int], list[float]]], +]: + tokenizer = ProjectedToolTokenizer() + exchanges = [] + messages: list[dict[str, Any]] = [] + records = [] + for index in range(2 if earlier_length else 1): + messages.append({"role": "user", "content": f"query{index}"}) + prompt = tokenizer.apply_chat_template(messages, add_generation_prompt=True) + terminal = not earlier_length or index == 1 + raw = ("raw tool output " * 8) if terminal else "earlier response\n\n" + message = ( + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "public_call", + "type": "function", + "function": {"name": "lookup", "arguments": "{}"}, + } + ], + } + if terminal + else {"role": "assistant", "content": raw} + ) + output = tokenizer._encode(raw + ("§" if terminal and sampled_eos else "")) + exchange = _chat_exchange(prompt, output, offset=index) + exchange.request["messages"] = cast( + list[ChatCompletionMessageParam], deepcopy(messages) + ) + payload = exchange.response.model_dump(mode="python") + payload["choices"][0]["message"] = message + payload["choices"][0]["finish_reason"] = finish if terminal else "length" + exchange.response = ChatCompletion.model_validate(payload) + exchanges.append(exchange) + logprobs = exchange.response.choices[0].logprobs + assert logprobs is not None and logprobs.content is not None + records.append((prompt, output, [entry.logprob for entry in logprobs.content])) + messages.append(message) + return ( + tr.Trajectory(exchanges=tr.TrajectoryExchanges(chat_completions=exchanges)), + tokenizer, + records, + ) + + +@pytest.mark.parametrize("finish", ["stop", "tool_calls", "length"]) +@pytest.mark.parametrize("sampled_eos", [False, True]) +@pytest.mark.parametrize("earlier_length", [False, True]) +def test_complete_native_terminal_does_not_reconstruct_projected_tool_body( + finish: str, + sampled_eos: bool, + earlier_length: bool, +) -> None: + trajectory, tokenizer, records = projected_history( + finish=finish, sampled_eos=sampled_eos, earlier_length=earlier_length + ) + original = trajectory.model_dump(mode="python") + result = trajectory.tokenize(tokenizer=tokenizer, multi_history=True) + assert len(result.histories) == 1 + history = result.histories[0] + assert history.tokens == records[-1][0] + records[-1][1] + required = ( + tr.TokenFlag.EXACT + | tr.TokenFlag.SAMPLED + | tr.TokenFlag.OUTPUT + | tr.TokenFlag.ASSISTANT + ) + for prompt, output, expected in records: + start, end = len(prompt), len(prompt) + len(output) + assert history.tokens[:end] == prompt + output + assert all(flag & required == required for flag in history.flags[start:end]) + assert history.logprobs[start:end] == expected + final_start = len(records[-1][0]) + sampled_stops = [ + i + for i in range(final_start, len(history.tokens)) + if history.flags[i] & tr.TokenFlag.STOP + ] + assert sampled_stops == ( + [len(history.tokens) - 1] if sampled_eos and finish != "length" else [] + ) + if earlier_length: + boundary = len(records[0][0]) + len(records[0][1]) + assert history.flags[boundary] == tr.TokenFlag.EXACT | tr.TokenFlag.STOP + assert math.isnan(history.logprobs[boundary]) + assert tr.first_occurrence_masks( + result.histories, where=tr.TokenFlag.OUTPUT + ) == tr.first_occurrence_masks(result.histories, where=tr.TokenFlag.SAMPLED) + assert trajectory.model_dump(mode="python") == original + + +def test_terminal_length_native_path_does_not_load_or_render( + monkeypatch: pytest.MonkeyPatch, +) -> None: + exchange = _chat_exchange([1], [2, 3]) + exchange.response.choices[0].finish_reason = "length" + trajectory = tr.Trajectory( + exchanges=tr.TrajectoryExchanges(chat_completions=[exchange]) + ) + + def unexpected(*args: Any, **kwargs: Any) -> Any: + raise AssertionError("complete terminal native data needs no renderer") + + monkeypatch.setattr(module, "_load_tokenizer", unexpected) + monkeypatch.setattr(module, "_tokenizer_config", unexpected) + result = trajectory.tokenize() + assert result.tokens == [1, 2, 3] + assert ( + result.flags[-1] + == tr.TokenFlag.EXACT + | tr.TokenFlag.SAMPLED + | tr.TokenFlag.OUTPUT + | tr.TokenFlag.ASSISTANT + ) + + +def test_explicit_template_override_retains_rendered_terminal_tail() -> None: + history, tokenizer, _ = _character_template_history(terminal_sampled_stop=False) + value = history.tokenize(tokenizer=tokenizer, chat_template="explicit renderer") + assert value.tokens[-1] == 9 + assert ( + value.flags[-1] + == tr.TokenFlag.ASSISTANT | tr.TokenFlag.OUTPUT | tr.TokenFlag.STOP + ) + + +def test_request_owned_assistant_roles_survive_terminal_native_tool_output() -> None: + trajectory, tokenizer, records = projected_history( + finish="tool_calls", sampled_eos=False, earlier_length=False + ) + exchange = trajectory.exchanges.chat_completions[0] + exchange.request["messages"].insert( + 0, {"role": "assistant", "content": "historical context"} + ) + prompt = tokenizer.apply_chat_template( + exchange.request["messages"], add_generation_prompt=True + ) + payload = exchange.response.model_dump(mode="python") + payload["prompt_token_ids"] = prompt + payload["choices"][0]["prompt_token_ids"] = prompt + exchange.response = ChatCompletion.model_validate(payload) + result = trajectory.tokenize(tokenizer=tokenizer) + output = records[-1][1] + assert result.tokens == prompt + output + prefix_flags = result.flags[: len(prompt)] + assert any(flag & tr.TokenFlag.ASSISTANT for flag in prefix_flags) + assert any(flag & tr.TokenFlag.STOP for flag in prefix_flags) + assert not any( + flag & (tr.TokenFlag.OUTPUT | tr.TokenFlag.SAMPLED) for flag in prefix_flags + ) + assert result.logprobs[len(prompt) :] == records[-1][2] + + +def test_unresolved_nonterminal_stop_still_loads_boundary_authority( + monkeypatch: pytest.MonkeyPatch, +) -> None: + trajectory, tokenizer, records = projected_history( + finish="length", sampled_eos=False, earlier_length=True + ) + trajectory.exchanges.chat_completions[0].response.choices[0].finish_reason = "stop" + loaded = [] + monkeypatch.setattr( + module, + "_tokenizer_config", + lambda *args: module._TokenizerConfig(base_model="test/model"), + ) + + def load(config: Any) -> Any: + loaded.append(config) + return tokenizer + + monkeypatch.setattr(module, "_load_tokenizer", load) + result = trajectory.tokenize() + assert len(loaded) == 1 + assert result.tokens == records[-1][0] + records[-1][1] + boundary = len(records[0][0]) + len(records[0][1]) + assert ( + result.flags[boundary] + == tr.TokenFlag.EXACT + | tr.TokenFlag.ASSISTANT + | tr.TokenFlag.OUTPUT + | tr.TokenFlag.STOP + ) + assert math.isnan(result.logprobs[boundary]) diff --git a/tests/unit/trajectories/test_recorded_boundaries.py b/tests/unit/trajectories/test_recorded_boundaries.py index 6cac07d9b..0e4af63c9 100644 --- a/tests/unit/trajectories/test_recorded_boundaries.py +++ b/tests/unit/trajectories/test_recorded_boundaries.py @@ -200,7 +200,9 @@ def apply_chat_template( flag & TokenFlag.SAMPLED for flag in tokenized.flags[len(prompt) : len(prompt) + len(output)] ) - assert sum(bool(flag & TokenFlag.STOP) for flag in tokenized.flags) == 2 + assert sum(bool(flag & TokenFlag.STOP) for flag in tokenized.flags) == 2 - ( + tool_position == 1 + ) assert trajectory.model_dump(mode="python") == before monkeypatch.setattr( module, "_tokenize_recorded_chat_boundaries", lambda *args, **kwargs: None @@ -217,12 +219,13 @@ def apply_chat_template( prompt, _ = expected_spans[1] assert old.tokens[: len(prompt)] != prompt else: - # The template owns this EOS; it must not become a sampled token. - assert old.tokens == tokenized.tokens[:-1] - tool_prompt, tool_output = expected_spans[tool_position] - stop_position = len(tool_prompt) + len(tool_output) - assert tokenized.flags[stop_position] & TokenFlag.STOP - assert not tokenized.flags[stop_position] & TokenFlag.SAMPLED + # A complete terminal native output owns the end of the history. + assert old.tokens == tokenized.tokens + if tool_position == 0: + tool_prompt, tool_output = expected_spans[tool_position] + stop_position = len(tool_prompt) + len(tool_output) + assert tokenized.flags[stop_position] & TokenFlag.STOP + assert not tokenized.flags[stop_position] & TokenFlag.SAMPLED @pytest.mark.parametrize("footer", ["footer§", "user-owned footer"]) @@ -365,7 +368,7 @@ def apply_chat_template( result = trajectory.tokenize(tokenizer=tokenizer, multi_history=True) assert len(result.histories) == 2 original, copied = result.histories - assert original.tokens == tokenizer._encode("turn 0ranswer§") + assert original.tokens == tokenizer._encode("turn 0ranswer") assert copied.tokens == tokenizer._encode("turn 0answer§turn 1answer§") copy_start, copy_end = len(prompt), len(prompt) + len("answer") assert copied.flags[copy_start:copy_end] == [ diff --git a/tests/unit/trajectories/test_tokenize.py b/tests/unit/trajectories/test_tokenize.py index 20b730996..27e438e13 100644 --- a/tests/unit/trajectories/test_tokenize.py +++ b/tests/unit/trajectories/test_tokenize.py @@ -413,7 +413,7 @@ def test_exact_sampled_tool_stop_is_stop_when_tokenizer_identifies_it() -> None: assert tokenized.flags[-1] == (_SAMPLED_ASSISTANT_OUTPUT | tr.TokenFlag.STOP) -def test_length_stop_keeps_sampled_content_and_adds_synthetic_stop() -> None: +def test_length_stop_ends_at_complete_native_output() -> None: exchange = _chat_exchange([1], [2]) exchange.response.choices[0].finish_reason = "length" @@ -421,11 +421,10 @@ def test_length_stop_keeps_sampled_content_and_adds_synthetic_stop() -> None: exchanges=TrajectoryExchanges(chat_completions=[exchange]) ).tokenize(tokenizer=_StopTokenizer()) - assert tokenized.tokens == [1, 2, 9] + assert tokenized.tokens == [1, 2] assert tokenized.flags == [ tr.TokenFlag.EXACT, _SAMPLED_ASSISTANT_OUTPUT, - tr.TokenFlag.STOP, ] @@ -547,7 +546,7 @@ def test_length_stop_mapping_allows_another_assistant_without_a_stop() -> None: assert tokenized.flags[-1] == tr.TokenFlag.STOP -def test_terminal_length_with_sampled_eos_still_adds_synthetic_stop() -> None: +def test_terminal_length_does_not_duplicate_or_relabel_sampled_eos() -> None: exchange = _chat_exchange([1], [2, 9]) exchange.response.choices[0].finish_reason = "length" @@ -555,10 +554,10 @@ def test_terminal_length_with_sampled_eos_still_adds_synthetic_stop() -> None: exchanges=TrajectoryExchanges(chat_completions=[exchange]) ).tokenize(tokenizer=_StopTokenizer()) - assert tokenized.tokens == [1, 2, 9, 9] + assert tokenized.tokens == [1, 2, 9] assert tokenized.flags[-2:] == [ _SAMPLED_ASSISTANT_OUTPUT, - tr.TokenFlag.STOP, + _SAMPLED_ASSISTANT_OUTPUT, ] @@ -928,7 +927,7 @@ def test_public_exact_chain_preserves_raw_drift_across_proven_length_boundary() @pytest.mark.parametrize("finish_reason", ["stop", "tool_calls"]) -def test_length_chain_retains_exact_prefix_with_terminal_synthetic_stop( +def test_length_chain_retains_exact_prefix_without_terminal_footer( finish_reason: Literal["stop", "tool_calls"], ) -> None: history, tokenizer, captured = _character_template_history( @@ -941,11 +940,9 @@ def test_length_chain_retains_exact_prefix_with_terminal_synthetic_stop( tokenized = history.tokenize(tokenizer=tokenizer) - assert tokenized.tokens == [*captured, 9] - assert all(flag & tr.TokenFlag.EXACT for flag in tokenized.flags[:-1]) - assert tokenized.flags[-1] == ( - tr.TokenFlag.STOP | tr.TokenFlag.ASSISTANT | tr.TokenFlag.OUTPUT - ) + assert tokenized.tokens == captured + assert all(flag & tr.TokenFlag.EXACT for flag in tokenized.flags) + assert tokenized.flags[-1] == (_SAMPLED_ASSISTANT_OUTPUT) assert sum(bool(flag & tr.TokenFlag.SAMPLED) for flag in tokenized.flags) == 19 @@ -1038,15 +1035,12 @@ def apply_chat_template( if mismatch: assert tokenized.tokens != expected return - assert tokenized.tokens == expected + assert tokenized.tokens == [*next_prompt, *tool_output] assert all( flag & tr.TokenFlag.EXACT for flag in tokenized.flags[: len(next_prompt) + len(tool_output)] ) - assert ( - tokenized.flags[-1] - == tr.TokenFlag.ASSISTANT | tr.TokenFlag.OUTPUT | tr.TokenFlag.STOP - ) + assert tokenized.flags[-1] == _SAMPLED_ASSISTANT_OUTPUT sampled = [ index for index, flag in enumerate(tokenized.flags) @@ -1487,7 +1481,7 @@ def test_metadata_only_final_token_is_preserved_as_sampled_stop() -> None: assert tokenized.flags[-1] == (_SAMPLED_ASSISTANT_OUTPUT | tr.TokenFlag.STOP) -def test_empty_output_materializes_a_synthetic_stop() -> None: +def test_recorded_empty_output_does_not_invent_a_response_token() -> None: exchange = _chat_exchange([1], []) exchange.response.choices[0].message.content = "" @@ -1495,10 +1489,9 @@ def test_empty_output_materializes_a_synthetic_stop() -> None: exchanges=TrajectoryExchanges(chat_completions=[exchange]) ).tokenize(tokenizer=_StopTokenizer()) - assert tokenized.tokens == [1, 9] + assert tokenized.tokens == [1] assert tokenized.flags == [ tr.TokenFlag.EXACT, - tr.TokenFlag.ASSISTANT | tr.TokenFlag.OUTPUT | tr.TokenFlag.STOP, ] From c5a30b799f4772c94148f9447dbc322953209f2c Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 06:39:28 +0000 Subject: [PATCH 14/16] Keep resolved STOP tests on histories that require role rendering --- .../test_resolved_stop_authority.py | 50 ++++++++++++++----- 1 file changed, 38 insertions(+), 12 deletions(-) diff --git a/tests/unit/trajectories/test_resolved_stop_authority.py b/tests/unit/trajectories/test_resolved_stop_authority.py index 17d15af81..e9fc6318a 100644 --- a/tests/unit/trajectories/test_resolved_stop_authority.py +++ b/tests/unit/trajectories/test_resolved_stop_authority.py @@ -16,12 +16,19 @@ def branch(index: int, *, length: bool, model: str = "test/model", eos: int = 9) output = _CharacterTemplateTokenizer._encode("answer") if not length: output.append(eos) - exchange = _chat_exchange( - _CharacterTemplateTokenizer._encode(text), output, model=model, offset=index - ) - exchange.request["messages"] = cast( - list[ChatCompletionMessageParam], [{"role": "user", "content": text}] - ) + messages: list[dict[str, str]] = [{"role": "user", "content": text}] + prompt = _CharacterTemplateTokenizer._encode(text) + if length: + # Loading is needed to prove request-owned assistant roles, independently + # of the final length stop, which now ends at its recorded output. + messages.insert(0, {"role": "assistant", "content": "historical context"}) + prompt = [ + *_CharacterTemplateTokenizer._encode("historical context"), + eos, + *prompt, + ] + exchange = _chat_exchange(prompt, output, model=model, offset=index) + exchange.request["messages"] = cast(list[ChatCompletionMessageParam], messages) exchange.response.choices[0].finish_reason = "length" if length else "stop" return exchange @@ -138,19 +145,29 @@ def test_each_model_uses_its_own_resolved_tokenizer(monkeypatch): ) result = trajectory( branch(0, length=True, model="model/a"), - branch(1, length=True, model="model/b"), + branch(1, length=True, model="model/b", eos=8), branch(2, length=False, model="model/a"), branch(3, length=False, model="model/b", eos=8), ).tokenize(multi_history=True) assert loads == ["model/a", "model/b"] - assert all(h.flags[-1] & tr.TokenFlag.STOP for h in result.histories) + assert [bool(h.flags[-1] & tr.TokenFlag.STOP) for h in result.histories] == [ + False, + True, + False, + True, + ] assert [h.model for h in result.histories] == [ "model/a", "model/a", "model/b", "model/b", ] - assert [h.tokens[-1] for h in result.histories] == [9, 9, 8, 8] + assert [h.tokens[-1] for h in result.histories] == [ + ord("r") + 100, + 9, + ord("r") + 100, + 8, + ] def test_conflicting_resolved_tokenizers_do_not_authorize_another_history(monkeypatch): @@ -158,7 +175,7 @@ def test_conflicting_resolved_tokenizers_do_not_authorize_another_history(monkey tokenizers = iter([_CharacterTemplateTokenizer(), OtherTokenizer()]) monkeypatch.setattr(module, "_load_tokenizer", lambda config: next(tokenizers)) result = trajectory( - branch(0, length=True), branch(1, length=True), branch(2, length=False) + branch(0, length=True), branch(1, length=True, eos=8), branch(2, length=False) ).tokenize(multi_history=True) assert not result.histories[-1].flags[-1] & tr.TokenFlag.STOP @@ -287,7 +304,7 @@ def __call__(self, text, **kwargs): assert caught.value is failure -def test_stop_postpass_keeps_copied_context_and_synthetic_tail_roles(monkeypatch): +def test_stop_postpass_keeps_copied_context_and_historical_roles(monkeypatch): bind(monkeypatch, {"test/model": _CharacterTemplateTokenizer()}) first = _chat_exchange([1], [2, 9]) second = _chat_exchange([1, 9, 4], [5, 9], offset=1) @@ -302,7 +319,16 @@ def test_stop_postpass_keeps_copied_context_and_synthetic_tail_roles(monkeypatch == tr.TokenFlag.EXACT | tr.TokenFlag.ASSISTANT | tr.TokenFlag.OUTPUT ) assert math.isnan(copied.logprobs[1]) - assert length.flags[-1] == tr.TokenFlag.STOP + assert length.flags[-1] == ( + tr.TokenFlag.EXACT + | tr.TokenFlag.SAMPLED + | tr.TokenFlag.ASSISTANT + | tr.TokenFlag.OUTPUT + ) + assert any( + flag & tr.TokenFlag.STOP and not flag & tr.TokenFlag.SAMPLED + for flag in length.flags + ) monkeypatch.setattr(module, "_complete_resolved_sampled_stops", lambda *args: None) baseline = value.tokenize(multi_history=True) same_except_stop(result, baseline) From afabb3c9faa66f8ec22bde736a4e5078046588b9 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 07:00:54 +0000 Subject: [PATCH 15/16] Test literal rendering separately from complete native output --- .../trajectories/test_literal_thinking_off.py | 77 ++++++++++++------- 1 file changed, 49 insertions(+), 28 deletions(-) diff --git a/tests/unit/trajectories/test_literal_thinking_off.py b/tests/unit/trajectories/test_literal_thinking_off.py index cabfce4b6..f956c0b09 100644 --- a/tests/unit/trajectories/test_literal_thinking_off.py +++ b/tests/unit/trajectories/test_literal_thinking_off.py @@ -19,6 +19,7 @@ Path(__file__).parents[2] / "fixtures/qwen35_preserved_thinking.jinja" ).read_text() _LITERAL = "HEAD\nDISCARDED_PUBLIC_SEGMENT\n\nTAIL" +_RENDER_OVERRIDE = _TEMPLATE + "{# explicit rendering #}" class _TemplateTokenizer: @@ -126,10 +127,10 @@ def _history( def _outcome( - history: tr.ChatCompletionsHistory, tokenizer: _TemplateTokenizer + history: tr.ChatCompletionsHistory, tokenizer: _TemplateTokenizer, **kwargs: Any ) -> object: try: - value = history.tokenize(tokenizer=tokenizer) + value = history.tokenize(tokenizer=tokenizer, **kwargs) except ValueError as error: return type(error), str(error) return value.tokens, value.flags, [None if x != x else x for x in value.logprobs] @@ -138,23 +139,25 @@ def _outcome( @pytest.mark.parametrize( "content", [_LITERAL, "literal text", "πonetwoend"] ) +@pytest.mark.parametrize("rendered", [False, True]) def test_native_thinking_off_retains_literal_content( - content: str, monkeypatch: pytest.MonkeyPatch + content: str, rendered: bool, monkeypatch: pytest.MonkeyPatch ) -> None: history, tokenizer = _history(content=content) original = history.model_dump(mode="python") - # The pre-fix history path misrenders literal content even when later native - # token splicing can recover the terminal output. - with monkeypatch.context() as patch: - patch.setattr( - _tokenize, "chat_template_with_preserved_thinking", lambda value: value - ) - _outcome(history, tokenizer) - assert content not in tokenizer.rendered[0] - tokenizer.calls.clear() - tokenizer.rendered.clear() - tokenized = history.tokenize(tokenizer=tokenizer) - assert content in tokenizer.rendered[0] + # Explicit rendering still needs literal-content normalization. Complete + # native output needs no rendering, even for a length-limited response. + override = _RENDER_OVERRIDE if rendered else None + if rendered: + with monkeypatch.context() as patch: + patch.setattr( + _tokenize, "chat_template_with_preserved_thinking", lambda value: value + ) + _outcome(history, tokenizer, chat_template=override) + assert content not in tokenizer.rendered[0] + tokenizer.calls.clear() + tokenizer.rendered.clear() + tokenized = history.tokenize(tokenizer=tokenizer, chat_template=override) sampled = [ i for i, flag in enumerate(tokenized.flags) if flag & tr.TokenFlag.SAMPLED ] @@ -163,8 +166,19 @@ def test_native_thinking_off_retains_literal_content( required = tr.TokenFlag.EXACT | tr.TokenFlag.ASSISTANT | tr.TokenFlag.OUTPUT assert all(tokenized.flags[i] & required == required for i in sampled) assert not any(flag & tr.TokenFlag.STOP for flag in tokenized.flags) - assert tokenizer.calls[0][-1].get("reasoning_content", "") == "" - assert tokenizer.calls[0][-1]["content"] == content + if rendered: + assert content in tokenizer.rendered[0] + assert tokenizer.calls[0][-1].get("reasoning_content", "") == "" + assert tokenizer.calls[0][-1]["content"] == content + else: + assert not tokenizer.calls and not tokenizer.rendered + source = history.message_sources[-1] + assert source is not None and isinstance( + source.exchange, tr.ChatCompletionsExchange + ) + recorded = source.exchange.response.choices[0].model_extra + assert recorded is not None + assert tokenized.tokens == recorded["prompt_token_ids"] + recorded["token_ids"] assert history.model_dump(mode="python") == original @@ -221,8 +235,9 @@ def test_literal_content_is_not_inferred_from_source_thinking_mode(case: str) -> if case == "visible_only": cast(dict[str, Any], history.messages[-1]).pop("reasoning") original = history.model_dump(mode="python") - _outcome(history, tokenizer) - # Plain content stays literal independently of recorded/current thinking mode + _outcome(history, tokenizer, chat_template=_RENDER_OVERRIDE) + # Exercise rendering explicitly even when complete native output can bypass + # it. Plain content stays literal independently of recorded/current thinking mode # and whether the message has complete native token metadata. Structured # reasoning remains a separate field on the render copy. assert tokenizer.calls[0][-1]["content"] == _LITERAL @@ -238,10 +253,13 @@ def test_literal_content_is_not_inferred_from_source_thinking_mode(case: str) -> @pytest.mark.parametrize("field", ["reasoning_content", "reasoning"]) -def test_explicit_empty_reasoning_is_preserved(field: str) -> None: +@pytest.mark.parametrize("rendered", [False, True]) +def test_explicit_empty_reasoning_is_preserved(field: str, rendered: bool) -> None: history, tokenizer = _history(reasoning="", reasoning_field=field) original = history.model_dump(mode="python") - tokenized = history.tokenize(tokenizer=tokenizer) + tokenized = history.tokenize( + tokenizer=tokenizer, chat_template=_RENDER_OVERRIDE if rendered else None + ) assert ( "".join( chr(token) @@ -250,7 +268,10 @@ def test_explicit_empty_reasoning_is_preserved(field: str) -> None: ) == _LITERAL ) - assert tokenizer.calls[0][-1].get("reasoning_content", "") == "" + if rendered: + assert tokenizer.calls[0][-1].get("reasoning_content", "") == "" + else: + assert not tokenizer.calls and not tokenizer.rendered assert history.model_dump(mode="python") == original @@ -409,17 +430,17 @@ def observe(*args: Any, **kwargs: Any) -> tr.TokenizedHistory | None: ) _outcome(history, tokenizer) boundary, old_exact = observed[0] - stored = list(boundary.tail + boundary.following) - assert old_exact is None - assert len(native_boundary) - len(stored) == 2 - assert stored[:-1] == native_boundary[:-3] - assert tokenizer.decode(stored[-1:]) == "\n" - assert tokenizer.decode(native_boundary[-3:]) == "\n\n\n\n" + # The final recorded body no longer needs a reconstructed terminal tail. + # Disabling literal normalization cannot invalidate the proved earlier gap. + assert old_exact is not None + assert list(boundary.tail + boundary.following) == native_boundary observed.clear() value = history.tokenize(tokenizer=tokenizer) fixed_boundary, fixed_exact = observed[0] assert fixed_exact is value + assert value.tokens == old_exact.tokens + assert value.flags == old_exact.flags assert list(fixed_boundary.tail + fixed_boundary.following) == native_boundary assert ( value.tokens[: len(last["prompt_token_ids"]) + len(last["token_ids"])] From d252b531571f1c4450e568a19a493e03d6cac91f Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 19:06:08 +0000 Subject: [PATCH 16/16] Preserve request roles and validate token evidence across callbacks --- docs/features/additional-histories.mdx | 13 +- src/art/trajectories/_tokenize.py | 180 +++++++++++++++--- src/art_inference/chat_template.py | 10 +- tests/unit/test_literal_reasoning_content.py | 29 +++ .../unit/trajectories/test_evidence_reuse.py | 169 ++++++++++++++++ .../unit/trajectories/test_native_terminal.py | 23 ++- .../test_resolved_stop_authority.py | 54 ++++++ tests/unit/trajectories/test_tokenize.py | 11 +- 8 files changed, 452 insertions(+), 37 deletions(-) diff --git a/docs/features/additional-histories.mdx b/docs/features/additional-histories.mdx index 154c9fd90..2a31c11bf 100644 --- a/docs/features/additional-histories.mdx +++ b/docs/features/additional-histories.mdx @@ -188,7 +188,18 @@ the original request and template reproduce its complete native prompt, ART can recover historical assistant roles from that rendering. This uses the original tool serialization order and preserves role labels through exact length-stop assembly; it does not restore destructive parsing for new response content. -Unproved historical role mappings still use the strict existing fallback. +With a supplied renderer, request-owned assistant roles are proved throughout +the native stream, including between sampled responses. A request whose text +disagrees with its recorded assistant tokens can still use the offline native +path, but cannot claim a complete rendered role mask. Unproved role mappings +use the strict existing fallback and may be refused. + +Custom STOP encoders and terminator decoders must not change already-consumed +source tokens, logprobs, model or stop evidence. ART checks each source around its +STOP callback and checks all consumed evidence before returning. Multi-history +tokenization also checks completed histories after later renderer callbacks. +These checks refuse stale results; they do not make callbacks or source objects +immutable. Callback-free native assembly retains its bounded evidence reuse. A response copied into a later, shortened prompt is output provenance, but it is not a fresh sample under that new prompt. ART keeps its `OUTPUT`, `ASSISTANT`, diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index 61626c6fc..051291739 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -2,7 +2,7 @@ from bisect import bisect_left import codecs -from collections.abc import Mapping, Sequence +from collections.abc import Callable, Mapping, Sequence from copy import deepcopy from dataclasses import dataclass, replace from datetime import datetime @@ -746,6 +746,34 @@ def _rendered_flag(assistant: bool, output: bool, stop: bool) -> TokenFlag: return flag | TokenFlag.STOP if stop else flag +def _merge_recorded_request_roles( + exact: TokenizedHistory, + rendered: Sequence[int], + assistant_mask: Sequence[bool], + output_mask: Sequence[bool], + stop_mask: Sequence[bool], + length_stop_mask: Sequence[bool], +) -> bool: + # Sampled responses own their native flags. Request-only assistant roles + # must retain a complete prefix proof, including roles between responses. + roles = [ + (index, _rendered_flag(assistant, False, stop)) + for index, (assistant, output, stop, length_stop) in enumerate( + zip(assistant_mask, output_mask, stop_mask, length_stop_mask, strict=True) + ) + if not output and not length_stop and (assistant or stop) + ] + end = roles[-1][0] + 1 if roles else 0 + if exact.tokens[:end] != list(rendered[:end]) or any( + exact.flags[index] & (TokenFlag.SAMPLED | TokenFlag.OUTPUT) + for index, _ in roles + ): + return False + for index, flags in roles: + exact.flags[index] |= flags + return True + + def _synthetic_length_stop_mask( messages: Sequence[Mapping[str, object]], sources: Sequence[object | None], @@ -915,6 +943,7 @@ class _TraceBuilder: trace: _HistoryTokenizationTrace | None = None tokenizer: Tokenizer | None = None rendered_outputs: tuple[tuple[int, int, object], ...] = () + validate_sources: Callable[[_SampledSourceKey | None], None] | None = None def set( self, @@ -930,6 +959,7 @@ def set( trace.validate(tokenized) self.trace = trace self.rendered_outputs = rendered_outputs + self.validate_sources = _sampled_source_validator(sources) def _fingerprint(value: object) -> str: @@ -4264,6 +4294,49 @@ def _source_has_no_materialized_output( return False +def _sampled_source_validator( + sources: Mapping[_SampledSourceKey, object], +) -> Callable[[_SampledSourceKey | None], None]: + expected = {} + for key, source in sources.items(): + exchange = _source_exchange(source) + if exchange is None: + raise ValueError("Sampled source has no exchange") + expected[key] = ( + source, + exchange, + exchange.model, + _source_stop_evidence(source, key), + ) + + def validate(selected: _SampledSourceKey | None) -> None: + for key in expected if selected is None else (selected,): + source, exchange, model, stop = expected[key] + current = ( + _exchange_sampled_source_key(source) + if isinstance(source, Exchange) + else _sampled_source_key(source) + ) + if ( + current != key + or _source_exchange(source) is not exchange + or exchange.model != model + or _source_stop_evidence(source, key) != stop + ): + raise ValueError("Sampled source changed during tokenization callback") + + return validate + + +def _stop_uses_callback(reason: int | str | None, tokenizer: Tokenizer | None) -> bool: + return tokenizer is not None and ( + isinstance(reason, str) + and bool(reason) + or not (isinstance(reason, int) and not isinstance(reason, bool)) + and callable(getattr(tokenizer, "convert_tokens_to_ids", None)) + ) + + def _mark_sampled_stops( token_ids: Sequence[int], flags: list[TokenFlag], @@ -4272,13 +4345,17 @@ def _mark_sampled_stops( *, tokenizer: Tokenizer | None, ) -> None: + validate_sources = None positions: dict[_SampledSourceKey, list[int]] = {} for index, source_key in enumerate(source_keys): if source_key is not None: positions.setdefault(source_key, []).append(index) for source_key, indices in positions.items(): + if validate_sources is not None: + validate_sources(source_key) source = sources[source_key] - if _source_stop_evidence(source, source_key)[0] != "stop": + kind, reason = _source_stop_evidence(source, source_key) + if kind != "stop": continue selected = [token_ids[index] for index in indices] complete = _source_output_tokens(source, source_key) @@ -4286,14 +4363,24 @@ def _mark_sampled_stops( continue if selected != complete[-len(selected) :]: continue + callback_used = _stop_uses_callback(reason, tokenizer) + if callback_used and validate_sources is None: + validate_sources = _sampled_source_validator(sources) count = _sampled_stop_suffix( selected, source=source, source_key=source_key, tokenizer=tokenizer, ) + if callback_used: + assert validate_sources is not None + validate_sources(source_key) for index in indices[-count:] if count else (): flags[index] |= TokenFlag.STOP + if validate_sources is not None: + # Later callbacks may edit an already-marked source. Check all consumed + # evidence once before return, without rehashing every source per stop. + validate_sources(None) @dataclass(frozen=True) @@ -4700,9 +4787,17 @@ def record( logprobs[start:end] = [math.nan] * len(retained_ids) records.clear() # A custom STOP decoder may change source objects. fingerprints.clear() + reason = _source_stop_evidence(source, source_key)[1] + validate_sources = ( + _sampled_source_validator({**sources, source_key: source}) + if _stop_uses_callback(reason, tokenizer) + else None + ) stop_count = _sampled_stop_suffix( output, source=source, source_key=source_key, tokenizer=tokenizer ) + if validate_sources is not None: + validate_sources(None) for offset in range(max(start, end - stop_count), end): flags[offset] |= TokenFlag.STOP boundary = (length_stop_boundaries or {}).get(source_key) @@ -4747,7 +4842,12 @@ def record( ): records.clear() # Never reuse records across user callbacks. fingerprints.clear() - if decode(native_boundary[:extra]).isspace(): + validate_sources = _sampled_source_validator( + {**sources, source_key: source} + ) + whitespace = decode(native_boundary[:extra]).isspace() + validate_sources(None) + if whitespace: # Services may insert whitespace before a truncated turn's # proven stop tail. Keep those served, nonsampled tokens. boundary = _RenderedLengthStopBoundary( @@ -6266,24 +6366,15 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: _trace=_trace, ) ) + and _merge_recorded_request_roles( + exact, + rendered, + assistant_mask, + output_mask, + stop_mask, + length_stop_mask, + ) ): - # The exact builder owns sampled spans, but the rendered request - # already proved role labels before the first sampled response. - # Retain those labels instead of discarding them on this return. - if ( - exact_prefix_length - and exact.tokens[:exact_prefix_length] == rendered[:exact_prefix_length] - and not any( - flag & TokenFlag.SAMPLED - for flag in exact.flags[:exact_prefix_length] - ) - ): - for index in range(exact_prefix_length): - exact.flags[index] |= _rendered_flag( - assistant_mask[index] and not length_stop_mask[index], - output_mask[index] and not length_stop_mask[index], - stop_mask[index], - ) return exact sampled_message_count = sum( @@ -7384,6 +7475,14 @@ def _tokenize_history( _trace=_trace, ) _validate_history_sources(history) + can_render = tokenizer is None or callable( + getattr(tokenizer, "apply_chat_template", None) + ) + needs_synthetic_stop = _history_needs_synthetic_stop(history, tokenizer) + if tokenizer is not None and can_render: + # STOP discovery can call a supplied tokenizer before native assembly. + # Its result cannot lend an old model/view proof to changed sources. + _validate_history_sources(history) override_requires_render = ( chat_template is not None and chat_template != getattr(history, "chat_template", None) @@ -7397,9 +7496,6 @@ def _tokenize_history( if _projection_validated else _history_render_state(history) ) - can_render = tokenizer is None or callable( - getattr(tokenizer, "apply_chat_template", None) - ) # Without a tokenizer, the stop decision and first exact assembly have no # user callback between them. Keep their evidence in one bounded phase; # never carry it into a rendered or tokenizer-supplied path. @@ -7411,7 +7507,6 @@ def _tokenize_history( has_length_stop = can_render and _history_has_length_stop( history, _fingerprints=fingerprints ) - needs_synthetic_stop = _history_needs_synthetic_stop(history, tokenizer) needs_render = ( render_state.needs_render or override_requires_render @@ -7428,6 +7523,17 @@ def _tokenize_history( ): return exact if isinstance(history, ChatCompletionsHistory): + needs_request_roles = ( + tokenizer is not None + and can_render + and any( + message.get("role") == "assistant" + and (source is None or not _source_is_sampled(source)) + for message, source in zip( + history.messages, history.message_sources, strict=True + ) + ) + ) if ( ( not has_length_stop @@ -7445,6 +7551,7 @@ def _tokenize_history( ) ) ) + and not needs_request_roles and not needs_synthetic_stop and not override_requires_render and not render_state.context_changed @@ -7475,7 +7582,7 @@ def _tokenize_history( ), _prior=_prior, _recorded_boundaries=( - (has_length_stop or needs_synthetic_stop) + (has_length_stop or needs_synthetic_stop or needs_request_roles) and not override_requires_render and not render_state.context_changed and (_projection_validated or render_state.projection_matches is True) @@ -7657,6 +7764,17 @@ def _materialize_trajectory( ) +def _validate_completed_sources(builders: Sequence[_TraceBuilder | None]) -> None: + if any( + builder is not None and builder.tokenizer is not None for builder in builders + ): + # Later callbacks may edit an earlier completed history. Check its + # original source keys and stop evidence without calling a tokenizer. + for builder in builders: + if builder is not None and builder.validate_sources is not None: + builder.validate_sources(None) + + def _complete_resolved_sampled_stops( tokenized: Sequence[TokenizedHistory], builders: Sequence[_TraceBuilder | None] ) -> None: @@ -7678,6 +7796,8 @@ def _complete_resolved_sampled_stops( and builder.trace is not None and (tokenizer := resolved.get(value.model)) is not None ): + assert builder.validate_sources is not None + builder.validate_sources(None) _mark_sampled_stops( value.tokens, value.flags, @@ -7685,6 +7805,7 @@ def _complete_resolved_sampled_stops( builder.trace.sources, tokenizer=tokenizer, ) + _validate_completed_sources(builders) def tokenize_trajectory( @@ -7721,9 +7842,8 @@ def tokenize_trajectory( prior: list[tuple[TokenizedHistory, _HistoryTokenizationTrace]] = [] tokenized = [] stop_builders: list[_TraceBuilder | None] = [] - collect_stops = tokenizer is None and len(histories) > 1 for history, copied in zip(histories, context_sources, strict=True): - trace = _TraceBuilder() if track_context or collect_stops else None + trace = _TraceBuilder() if len(histories) > 1 else None result = tokenize_history( history, model=model if isinstance(history, LegacyHistory) else history.model, @@ -7740,8 +7860,7 @@ def tokenize_trajectory( stop_builders.append(trace) if track_context and trace is not None and trace.trace is not None: prior.append((result, trace.trace)) - if collect_stops: - _complete_resolved_sampled_stops(tokenized, stop_builders) + _complete_resolved_sampled_stops(tokenized, stop_builders) if not multi_history: return _materialize_trajectory(tokenized[0], trajectory) return TokenizedMultiHistoryTrajectory( @@ -7790,8 +7909,7 @@ def _tokenize_trajectory_with_trace( tokenized_histories.append(tokenized) traces.append(trace_builder.trace) builders.append(trace_builder) - if tokenizer is None: - _complete_resolved_sampled_stops(tokenized_histories, builders) + _complete_resolved_sampled_stops(tokenized_histories, builders) return ( TokenizedMultiHistoryTrajectory( trajectory=trajectory, diff --git a/src/art_inference/chat_template.py b/src/art_inference/chat_template.py index bf54218ee..93b18f294 100644 --- a/src/art_inference/chat_template.py +++ b/src/art_inference/chat_template.py @@ -104,7 +104,15 @@ def operations(text: str): end = selected[-1][3] try: if operations(template[start:end]) == operation: - edits[start, end] = "" + # The enclosing tags also control unrelated surrounding + # whitespace. Disable the parser without deleting those tags. + _, body_start, body_end, first_end = selected[0] + edits[start, end] = ( + template[start:body_start] + + " if false " + + template[body_end:first_end] + + template[selected[-1][0] : end] + ) except TemplateSyntaxError: continue if not edits: diff --git a/tests/unit/test_literal_reasoning_content.py b/tests/unit/test_literal_reasoning_content.py index dff7da9f2..1c8056ec8 100644 --- a/tests/unit/test_literal_reasoning_content.py +++ b/tests/unit/test_literal_reasoning_content.py @@ -34,6 +34,35 @@ ) +@pytest.mark.parametrize("trim_blocks,lstrip_blocks", [(False, False), (True, True)]) +@pytest.mark.parametrize( + "left,right", [("", ""), ("-", ""), ("", "-"), ("-", "-"), ("+", "+")] +) +@pytest.mark.parametrize("newline", ["\n", "\r\n", "\r"]) +def test_disabling_inline_parser_preserves_outer_whitespace( + trim_blocks, lstrip_blocks, left, right, newline +): + match = _QWEN_INLINE_REASONING.search(_TEMPLATE) + assert match is not None + operation = match.group() + operation = "{%" + left + operation[3:] + operation = operation[:-2].rstrip("-+") + right + "%}" + template = ("HEADER \n\t" + operation + "\n \tTAIL{{ content }}").replace( + "\n", newline + ) + env = ImmutableSandboxedEnvironment( + trim_blocks=trim_blocks, lstrip_blocks=lstrip_blocks + ) + fixed = _without_inline_reasoning_parser(template) + ordinary = env.from_string(template).render(content="plain answer") + assert env.from_string(fixed).render(content="plain answer") == ordinary + literal = "prefixliteralsuffix" + assert env.from_string(fixed).render(content=literal) == ordinary.replace( + "plain answer", literal + ) + assert _without_inline_reasoning_parser(fixed) == fixed + + def _render(template, messages, **kwargs): def refuse(message): raise ValueError(message) diff --git a/tests/unit/trajectories/test_evidence_reuse.py b/tests/unit/trajectories/test_evidence_reuse.py index 8555ffe96..31ff8d55c 100644 --- a/tests/unit/trajectories/test_evidence_reuse.py +++ b/tests/unit/trajectories/test_evidence_reuse.py @@ -306,6 +306,175 @@ def __call__(self, text, **kwargs): assert not value.flags[1] & tr.TokenFlag.SAMPLED +@pytest.mark.parametrize( + "mutation_call,field", + [(call, "logprob") for call in range(1, 7)] + + [(2, field) for field in ("prompt", "output", "finish", "stop_reason", "model")], +) +def test_public_copied_stop_callback_cannot_return_stale_final_logprobs( + mutation_call, field +): + first = _chat_exchange([1], [2, 3]) + extras(first)["stop_reason"] = "public-stop" + second = _chat_exchange([1, 3, 4], [5, 6], offset=1) + value = trajectory(first, second) + + class StopTokenizer: + eos_token_id = 6 + all_special_ids = [] + calls = 0 + + def __call__(self, text, **kwargs): + assert text == "public-stop" + self.calls += 1 + if self.calls == mutation_call: + if field == "logprob": + logprobs(second)[0].logprob = -8.5 + elif field in {"prompt", "output"}: + extras(second)[ + "prompt_token_ids" if field == "prompt" else "token_ids" + ][0] = 99 + elif field == "finish": + second.response.choices[0].finish_reason = "length" + elif field == "stop_reason": + extras(second)["stop_reason"] = 99 + else: + second.request["model"] = "changed/model" + return {"input_ids": [3]} + + tokenizer = StopTokenizer() + if mutation_call == 2: + with pytest.raises( + ValueError, match="Sampled source changed during tokenization callback" + ): + value.tokenize(tokenizer=tokenizer, multi_history=True) + assert tokenizer.calls == 2 + return + result = value.tokenize(tokenizer=tokenizer, multi_history=True) + assert len(result.histories) == 2 + assert result.histories[-1].tokens == [1, 3, 4, 5, 6] + assert result.histories[-1].logprobs[3] == logprobs(second)[0].logprob + + +def test_final_stop_marker_callback_cannot_change_consumed_source(): + exchange = _chat_exchange([1], [2, 3]) + extras(exchange)["stop_reason"] = "public-stop" + + class StopTokenizer: + def __call__(self, text, **kwargs): + assert text == "public-stop" + extras(exchange)["stop_reason"] = 99 + return {"input_ids": [3]} + + with pytest.raises( + ValueError, match="Sampled source changed during tokenization callback" + ): + trajectory(exchange).tokenize(tokenizer=StopTokenizer()) + + +def test_stop_callback_is_checked_before_next_source_can_restore_evidence(): + first = _chat_exchange([1], [2, 3]) + second = _chat_exchange([1, 2, 3, 4], [5, 6], offset=1) + extras(first)["stop_reason"] = "first-stop" + extras(second)["stop_reason"] = "second-stop" + calls = [] + + class StopTokenizer: + def __call__(self, text, **kwargs): + calls.append(text) + if text == "first-stop": + extras(second)["stop_reason"] = "changed-stop" + return {"input_ids": [3]} + extras(second)["stop_reason"] = "second-stop" + return {"input_ids": [99]} + + with pytest.raises( + ValueError, match="Sampled source changed during tokenization callback" + ): + trajectory(first, second).tokenize(tokenizer=StopTokenizer()) + assert calls == ["first-stop"] + + +def test_later_stop_callback_cannot_change_an_already_marked_source(): + first = _chat_exchange([1], [2, 3]) + second = _chat_exchange([1, 2, 3, 4], [5, 6], offset=1) + extras(first)["stop_reason"] = "first-stop" + extras(second)["stop_reason"] = "second-stop" + + class StopTokenizer: + def __call__(self, text, **kwargs): + if text == "second-stop": + logprobs(first)[0].logprob = -9 + return {"input_ids": [3 if text == "first-stop" else 6]} + + with pytest.raises( + ValueError, match="Sampled source changed during tokenization callback" + ): + trajectory(first, second).tokenize(tokenizer=StopTokenizer()) + + +@pytest.mark.parametrize("callback", ["convert", "decode"]) +def test_terminator_lookup_cannot_change_consumed_logprobs(callback): + exchange = _chat_exchange([1], [2, 3]) + + class Tokenizer: + eos_token_id = 3 + unk_token_id = None + all_special_tokens = [] + + def convert_tokens_to_ids(self, token): + if callback == "convert": + logprobs(exchange)[0].logprob = -9 + return 99 + + def decode(self, ids, **kwargs): + if callback == "decode": + logprobs(exchange)[0].logprob = -9 + return "not a special token" + + with pytest.raises( + ValueError, match="Sampled source changed during tokenization callback" + ): + trajectory(exchange).tokenize(tokenizer=Tokenizer()) + + +@pytest.mark.parametrize("reason", [None, 3, ""]) +def test_plain_stop_authority_needs_no_callback_revalidation(monkeypatch, reason): + exchange = _chat_exchange([1], [2, 3]) + extras(exchange)["stop_reason"] = reason + + class Tokenizer: + eos_token_id = 3 + + def unexpected(*args): + raise AssertionError("plain attributes and numeric STOP need no callback fence") + + monkeypatch.setattr(module, "_sampled_source_validator", unexpected) + result = trajectory(exchange).tokenize(tokenizer=Tokenizer()) + assert result.tokens == [1, 2, 3] + assert result.logprobs[1:] == [-0.2, -0.3] + assert result.flags[-1] & tr.TokenFlag.STOP + + +def test_stop_decision_callback_cannot_change_history_model(): + exchange = _chat_exchange([1], [2, 3]) + extras(exchange)["stop_reason"] = "public-stop" + history = trajectory(exchange).histories()[0] + + class Tokenizer: + eos_token_id = 3 + + def apply_chat_template(self, *args, **kwargs): + raise AssertionError("model change must be refused before rendering") + + def __call__(self, text, **kwargs): + exchange.request["model"] = "changed/model" + return {"input_ids": [3]} + + with pytest.raises(ValueError, match="model no longer matches"): + history.tokenize(tokenizer=Tokenizer()) + + def test_decoder_exception_identity_and_next_call_fresh(): first = _chat_exchange([1], [2, 3]) extras(first)["stop_reason"] = "public-stop" diff --git a/tests/unit/trajectories/test_native_terminal.py b/tests/unit/trajectories/test_native_terminal.py index 7858806aa..da2339027 100644 --- a/tests/unit/trajectories/test_native_terminal.py +++ b/tests/unit/trajectories/test_native_terminal.py @@ -176,9 +176,13 @@ def test_explicit_template_override_retains_rendered_terminal_tail() -> None: ) -def test_request_owned_assistant_roles_survive_terminal_native_tool_output() -> None: +@pytest.mark.parametrize("finish,sampled_eos", [("tool_calls", False), ("stop", True)]) +@pytest.mark.parametrize("standalone", [False, True]) +def test_request_owned_assistant_roles_survive_terminal_native_tool_output( + finish: str, sampled_eos: bool, standalone: bool +) -> None: trajectory, tokenizer, records = projected_history( - finish="tool_calls", sampled_eos=False, earlier_length=False + finish=finish, sampled_eos=sampled_eos, earlier_length=False ) exchange = trajectory.exchanges.chat_completions[0] exchange.request["messages"].insert( @@ -191,7 +195,13 @@ def test_request_owned_assistant_roles_survive_terminal_native_tool_output() -> payload["prompt_token_ids"] = prompt payload["choices"][0]["prompt_token_ids"] = prompt exchange.response = ChatCompletion.model_validate(payload) - result = trajectory.tokenize(tokenizer=tokenizer) + original = trajectory.model_dump(mode="python") + if standalone: + selected = trajectory.histories()[0] + assert isinstance(selected, tr.ChatCompletionsHistory) + else: + selected = trajectory + result = selected.tokenize(tokenizer=tokenizer) output = records[-1][1] assert result.tokens == prompt + output prefix_flags = result.flags[: len(prompt)] @@ -201,6 +211,13 @@ def test_request_owned_assistant_roles_survive_terminal_native_tool_output() -> flag & (tr.TokenFlag.OUTPUT | tr.TokenFlag.SAMPLED) for flag in prefix_flags ) assert result.logprobs[len(prompt) :] == records[-1][2] + expected = [tr.TokenFlag.EXACT] * len(prompt) + start = len(tokenizer._encode("assistant:")) + end = len(tokenizer._encode("assistant:historical context§")) + expected[start:end] = [tr.TokenFlag.EXACT | tr.TokenFlag.ASSISTANT] * (end - start) + expected[end - 1] |= tr.TokenFlag.STOP + assert prefix_flags == expected + assert trajectory.model_dump(mode="python") == original def test_unresolved_nonterminal_stop_still_loads_boundary_authority( diff --git a/tests/unit/trajectories/test_resolved_stop_authority.py b/tests/unit/trajectories/test_resolved_stop_authority.py index e9fc6318a..fdfd2cf90 100644 --- a/tests/unit/trajectories/test_resolved_stop_authority.py +++ b/tests/unit/trajectories/test_resolved_stop_authority.py @@ -109,6 +109,60 @@ def test_all_exact_histories_do_not_load_or_retain_prior_call_authority(monkeypa assert all(not (h.flags[-1] & tr.TokenFlag.STOP) for h in result.histories) +@pytest.mark.parametrize("supplied", [False, True]) +def test_later_history_callback_cannot_change_completed_source(monkeypatch, supplied): + first = branch(0, length=False, model="model/a") + second = branch(1, length=True, model="model/b") + + class Tokenizer(_CharacterTemplateTokenizer): + def apply_chat_template(self, messages, **kwargs): + assert first.response.choices[0].logprobs is not None + assert first.response.choices[0].logprobs.content is not None + first.response.choices[0].logprobs.content[0].logprob = -99 + return super().apply_chat_template(messages, **kwargs) + + tokenizer = Tokenizer() + bind(monkeypatch, {"model/b": tokenizer}) + with pytest.raises( + ValueError, match="Sampled source changed during tokenization callback" + ): + trajectory(first, second).tokenize( + multi_history=True, tokenizer=tokenizer if supplied else None + ) + + +def test_stop_postpass_checks_original_evidence_before_consuming_it(monkeypatch): + first = branch(0, length=False) + second = branch(1, length=False) + last = branch(2, length=True) + first_extra = first.response.choices[0].model_extra + second_extra = second.response.choices[0].model_extra + assert first_extra is not None and second_extra is not None + second_extra["stop_reason"] = "restore" + + class Tokenizer(_CharacterTemplateTokenizer): + restores = 0 + + def apply_chat_template(self, messages, **kwargs): + first_extra["stop_reason"] = 999 + return super().apply_chat_template(messages, **kwargs) + + def __call__(self, text, **kwargs): + if text == "restore": + self.restores += 1 + first_extra.pop("stop_reason", None) + return {"input_ids": [9]} + return super().__call__(text, **kwargs) + + tokenizer = Tokenizer() + bind(monkeypatch, {"test/model": tokenizer}) + with pytest.raises( + ValueError, match="Sampled source changed during tokenization callback" + ): + trajectory(first, second, last).tokenize(multi_history=True) + assert tokenizer.restores == 0 + + def test_resolved_authority_never_crosses_model_identity(monkeypatch): loads = bind(monkeypatch, {"model/a": _CharacterTemplateTokenizer()}) result = trajectory( diff --git a/tests/unit/trajectories/test_tokenize.py b/tests/unit/trajectories/test_tokenize.py index 27e438e13..fc86d06c3 100644 --- a/tests/unit/trajectories/test_tokenize.py +++ b/tests/unit/trajectories/test_tokenize.py @@ -9093,7 +9093,16 @@ def apply_chat_template( _tokenize_trajectory_with_trace(trajectory, tokenizer=tokenizer) return - tokenized, traces = _tokenize_trajectory_with_trace(trajectory, tokenizer=tokenizer) + if corruption == "changed_sampled_token": + # The later request text disagrees with its recorded assistant token. + # A supplied renderer cannot prove that request's full role mask. + with pytest.raises(ValueError, match="Cannot preserve assistant boundaries"): + _tokenize_trajectory_with_trace(trajectory, tokenizer=tokenizer) + tokenized, traces = _tokenize_trajectory_with_trace(trajectory, tokenizer=None) + else: + tokenized, traces = _tokenize_trajectory_with_trace( + trajectory, tokenizer=tokenizer + ) assert len(tokenized.histories) == ( 2 if corruption == "changed_sampled_token" else 1 )