From dda11089a6549fb50421e3743f3a489e81bb709a Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 19:07:37 +0000 Subject: [PATCH 01/67] 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/67] 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/67] 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/67] 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/67] 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/67] 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/67] 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/67] 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/67] 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/67] 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/67] 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/67] 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/67] 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/67] 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/67] 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/67] 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 ) From 360f126858c4a452b54d69721256a742f507cfbf Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 19:43:35 +0000 Subject: [PATCH 17/67] Preserve complete tokenization context across renderer callbacks --- docs/features/additional-histories.mdx | 12 +- src/art/trajectories/_tokenize.py | 304 +++++++++++++++--- src/art_inference/chat_template.py | 180 +++++++++-- tests/unit/test_literal_reasoning_content.py | 199 ++++++++++++ .../unit/trajectories/test_evidence_reuse.py | 90 ++++++ .../test_recorded_prompt_roles.py | 92 ++++++ .../test_resolved_stop_authority.py | 186 +++++++++++ 7 files changed, 990 insertions(+), 73 deletions(-) diff --git a/docs/features/additional-histories.mdx b/docs/features/additional-histories.mdx index 2a31c11bf..f15043116 100644 --- a/docs/features/additional-histories.mdx +++ b/docs/features/additional-histories.mdx @@ -191,15 +191,21 @@ assembly; it does not restore destructive parsing for new response content. 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. +path, but cannot claim a complete rendered role mask. When literal-content +normalization changes historical rendering, ART requires the original request +to reproduce the complete recorded prompt before recovering those roles. If +that proof fails, ART refuses to attach recorded logprobs to a changed prompt. 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. +Ordered original requests and history context, including roles, tools, render +settings and protocol selectors, are captured before callbacks and checked +before reuse. Unexpected callback exceptions keep their original identity. These checks refuse stale results; they do not make callbacks or source objects -immutable. Callback-free native assembly retains its bounded evidence reuse. +immutable. Context snapshots reuse shared containers only within one observation; +callback-free native assembly retains its bounded response-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 051291739..0d9c0d4ba 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -944,6 +944,8 @@ class _TraceBuilder: tokenizer: Tokenizer | None = None rendered_outputs: tuple[tuple[int, int, object], ...] = () validate_sources: Callable[[_SampledSourceKey | None], None] | None = None + validate_context: Callable[[bool], None] | None = None + track_sources: bool = True def set( self, @@ -959,7 +961,8 @@ def set( trace.validate(tokenized) self.trace = trace self.rendered_outputs = rendered_outputs - self.validate_sources = _sampled_source_validator(sources) + if self.track_sources: + self.validate_sources = _sampled_source_validator(sources) def _fingerprint(value: object) -> str: @@ -4294,6 +4297,71 @@ def _source_has_no_materialized_output( return False +def _tokenization_context(value: object) -> object: + """Snapshot ordered semantic inputs without serializing sampled responses. + + History and source fields remain typed, including protocol selectors and + request-owned context. Sampled response evidence has its own source-key + validator; retaining entire response objects here would duplicate it. + """ + # Shared message dictionaries occur in many recorded request prefixes. + # Intern only within this one observation, never across callback checks. + observed: dict[int, tuple[object, object]] = {} + + def snapshot(item: object) -> object: + kind = type(item) + if kind in (str, int, bool, bytes, type(None), datetime): + return kind, item + if kind is float: + return kind, repr(item) + previous = observed.get(id(item)) + if previous is not None and previous[0] is item: + return previous[1] + if kind in (list, tuple): + result = kind, tuple(snapshot(child) for child in cast(Sequence, item)) + elif isinstance(item, Mapping): + result = ( + kind, + tuple((snapshot(key), snapshot(child)) for key, child in item.items()), + ) + elif isinstance(item, Exchange): + result = kind, id(item), item.model, snapshot(item.request) + elif isinstance(item, BaseModel): + result = ( + kind, + tuple( + (name, snapshot(getattr(item, name))) + for name in type(item).model_fields + ), + snapshot(item.model_extra), + ) + else: + raise TypeError("Unsupported mutable tokenization context") + observed[id(item)] = item, result + return result + + return snapshot(value) + + +def _tokenization_context_validator(value: object) -> Callable[[bool], None]: + try: + expected = _tokenization_context(value) + except TypeError: + # An opaque, complete native history still works without a tokenizer. + # A callback-bearing path cannot claim an uncheckable context proof. + expected = None + + def validate(require_supported: bool) -> None: + if expected is None and not require_supported: + return + if expected is None or _tokenization_context(value) != expected: + raise ValueError( + "Tokenization context changed during tokenization callback" + ) + + return validate + + def _sampled_source_validator( sources: Mapping[_SampledSourceKey, object], ) -> Callable[[_SampledSourceKey | None], None]: @@ -4307,11 +4375,12 @@ def _sampled_source_validator( exchange, exchange.model, _source_stop_evidence(source, key), + _tokenization_context_validator(exchange.request), ) def validate(selected: _SampledSourceKey | None) -> None: for key in expected if selected is None else (selected,): - source, exchange, model, stop = expected[key] + source, exchange, model, stop, validate_request = expected[key] current = ( _exchange_sampled_source_key(source) if isinstance(source, Exchange) @@ -4324,6 +4393,7 @@ def validate(selected: _SampledSourceKey | None) -> None: or _source_stop_evidence(source, key) != stop ): raise ValueError("Sampled source changed during tokenization callback") + validate_request(True) return validate @@ -4926,7 +4996,7 @@ def _tokenize_recorded_chat_boundaries( 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() + sources: dict[_SampledSourceKey, object] = {} for index, (message, source) in enumerate( zip(messages, history.message_sources, strict=True) ): @@ -4936,9 +5006,9 @@ def _tokenize_recorded_chat_boundaries( 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: + if key in sources or prompt is None or output is None: return None - seen.add(key) + sources[key] = source entries.append((index, source, prompt, output, logprobs)) if not entries: return None @@ -4966,73 +5036,120 @@ def _tokenize_recorded_chat_boundaries( ): return None boundaries: dict[_SampledSourceKey, _RenderedLengthStopBoundary] = {} - terminators = _terminator_ids(tokenizer) + # Entries already consumed every native record. No callback may replace + # that evidence before a later source or the final assembler reads it again. + validate_sources = _sampled_source_validator(sources) + validate_context = _tokenization_context_validator(history) + keys = tuple(sources) + selected_key: _SampledSourceKey | None = None + + def checked(call: Callable[[], Any], *, optional_decode: bool = False) -> Any: + try: + value = call() + except (TypeError, KeyError, NotImplementedError): + # These capability failures are caught by the caller and may + # resume tokenization. A changed source must not reach fallback. + validate_sources(selected_key) + validate_context(True) + raise + except ValueError: + if not optional_decode: + raise + value = None + # Other callback exceptions propagate unchanged; no result or fallback + # can consume their potentially changed inputs. + validate_sources(selected_key) + validate_context(True) + return value + + def decline() -> None: + validate_sources(None) + validate_context(True) + + terminators = checked(lambda: _terminator_ids(tokenizer)) if not terminators: - return None + return decline() # 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) + selected_key = key = keys[ordinal] + validate_sources(key) 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 + return decline() + if stop == "stop" and checked( + lambda: _sampled_stop_suffix( + output, source=source, source_key=key, tokenizer=tokenizer + ) ): continue try: - try: - body = decode( + body = checked( + lambda: 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) + ), + optional_decode=True, + ) + if body is None: + return decline() + generation = checked( + lambda: render(messages[:index], add_generation_prompt=True) + ) + completed = checked( + lambda: 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 + return decline() suffix = completed[len(generation) + len(body) :] - tail = _ids(tokenizer(suffix, add_special_tokens=False)) if suffix else [] + tail = ( + _ids(checked(lambda: 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 + return decline() terminator = stops[0] - try: - trailing = decode( + trailing = checked( + lambda: decode( tail[terminator + 1 :], skip_special_tokens=False, clean_up_tokenization_spaces=False, - ) - except ValueError: - return None + ), + optional_decode=True, + ) if not isinstance(trailing, str) or trailing and not trailing.isspace(): - return None + return decline() 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 + next_generation = checked( + lambda: render(messages[:next_index], add_generation_prompt=True) ) if not next_generation.startswith(completed): - return None + return decline() gap = suffix + next_generation[len(completed) :] - gap_ids = _ids(tokenizer(gap, add_special_tokens=False)) + gap_ids = _ids( + checked(lambda: tokenizer(gap, add_special_tokens=False)) + ) if ( gap_ids[: len(tail)] != tail or next_prompt[len(prompt) + len(output) :] != gap_ids ): - return None + return decline() 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 decline() + validate_sources(None) + validate_context(True) return _tokenize_exact_projected_chat_history( history, tokenizer=tokenizer, @@ -6366,16 +6483,80 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: _trace=_trace, ) ) - and _merge_recorded_request_roles( + ): + if _merge_recorded_request_roles( exact, rendered, assistant_mask, output_mask, stop_mask, length_stop_mask, - ) - ): - return exact + ): + return exact + if ( + _recorded_boundaries + and chat_template is None + and chat_template_kwargs is None + ): + # Later request-only assistants can have historical rendering + # different from today's normalized template too. Prove the + # entire final recorded request, not only the first prompt. + source = history.message_sources[sampled_message_indices[-1]] + assert source is not None + exchange = _source_exchange(source) + assert isinstance( + exchange, + (ChatCompletionsExchange, MessagesExchange, ResponsesExchange), + ) + prompt = source_prompt_tokens(source) + signature = _source_signature(source) + validate_context = _tokenization_context_validator(history) + masks = None + if prompt and exact.tokens[: len(prompt)] == prompt: + request_messages, request_tools = _request_messages(exchange) + try: + masks = _recorded_prompt_role_masks( + request_messages, + [None] * len(request_messages), + prompt, + tokenizer=resolved_tokenizer, + template=original_template, + tools=request_tools, + kwargs=kwargs, + ) + except (TypeError, KeyError, NotImplementedError): + pass + validate_context(True) + prompt_cache.clear() + output_cache.clear() + if _source_signature(source) != signature: + raise ValueError( + "Sampled source changed while proving recorded request roles" + ) + if ( + masks is not None + and all( + flag & (TokenFlag.SAMPLED | TokenFlag.OUTPUT) + or not (flag & TokenFlag.ASSISTANT) + or assistant + for flag, assistant in zip(exact.flags, masks[0]) + ) + and all( + flag & (TokenFlag.SAMPLED | TokenFlag.OUTPUT) + or not (flag & TokenFlag.STOP) + or stop + for flag, stop in zip(exact.flags, masks[1]) + ) + ): + for index, (assistant, stop) in enumerate(zip(*masks, strict=True)): + if not exact.flags[index] & ( + TokenFlag.SAMPLED | TokenFlag.OUTPUT + ): + exact.flags[index] |= _rendered_flag(assistant, False, stop) + return exact + raise ValueError( + "Cannot preserve request roles across exact native prompt replacement" + ) sampled_message_count = sum( message.get("role") == "assistant" @@ -7668,12 +7849,20 @@ def tokenize_history( _projection_validated: bool = False, _context_sources: Sequence[object] | None = None, ) -> TokenizedHistory: + trace_builder = _trace or _TraceBuilder(track_sources=False) + if trace_builder.validate_context is None: + trace_builder.validate_context = _tokenization_context_validator( + [history, chat_template, chat_template_kwargs] + ) + else: + trace_builder.validate_context(False) copied = ( list(_context_sources) if _context_sources is not None else _partial_native_context(history) ) if copied: + trace_builder.track_sources = True history = cast(History, history) _validate_history_sources(history) state = None if _projection_validated else _history_render_state(history) @@ -7708,7 +7897,6 @@ def tokenize_history( 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, @@ -7721,6 +7909,8 @@ def tokenize_history( _projection_validated=_projection_validated, _copied_context=bool(copied), ) + if trace_builder.tokenizer is not None: + trace_builder.validate_context(True) if copied: if trace_builder is None or trace_builder.trace is None: raise ValueError( @@ -7771,8 +7961,11 @@ def _validate_completed_sources(builders: Sequence[_TraceBuilder | None]) -> Non # 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) + if builder is not None: + if builder.validate_context is not None: + builder.validate_context(True) + if builder.validate_sources is not None: + builder.validate_sources(None) def _complete_resolved_sampled_stops( @@ -7797,6 +7990,8 @@ def _complete_resolved_sampled_stops( and (tokenizer := resolved.get(value.model)) is not None ): assert builder.validate_sources is not None + if builder.validate_context is not None: + builder.validate_context(True) builder.validate_sources(None) _mark_sampled_stops( value.tokens, @@ -7841,9 +8036,18 @@ def tokenize_trajectory( track_context = len(histories) > 1 and any(context_sources) prior: list[tuple[TokenizedHistory, _HistoryTokenizationTrace]] = [] tokenized = [] - stop_builders: list[_TraceBuilder | None] = [] - for history, copied in zip(histories, context_sources, strict=True): - trace = _TraceBuilder() if len(histories) > 1 else None + stop_builders = [ + _TraceBuilder( + track_sources=len(histories) > 1, + validate_context=_tokenization_context_validator( + [history, chat_template, chat_template_kwargs] + ), + ) + for history in histories + ] + for history, copied, trace in zip( + histories, context_sources, stop_builders, strict=True + ): result = tokenize_history( history, model=model if isinstance(history, LegacyHistory) else history.model, @@ -7857,7 +8061,6 @@ def tokenize_trajectory( _context_sources=copied, ) tokenized.append(result) - stop_builders.append(trace) if track_context and trace is not None and trace.trace is not None: prior.append((result, trace.trace)) _complete_resolved_sampled_stops(tokenized, stop_builders) @@ -7886,13 +8089,19 @@ def _tokenize_trajectory_with_trace( histories = trajectory.histories(model=model) tokenized_histories: list[TokenizedHistory] = [] traces: list[_HistoryTokenizationTrace] = [] - builders: list[_TraceBuilder] = [] - for history in histories: + builders = [ + _TraceBuilder( + validate_context=_tokenization_context_validator( + [history, chat_template, chat_template_kwargs] + ) + ) + for history in histories + ] + for history, trace_builder in zip(histories, builders, strict=True): if isinstance(history, LegacyHistory): raise AssertionError( "Exchange trajectories cannot produce legacy histories" ) - trace_builder = _TraceBuilder() tokenized = tokenize_history( history, model=history.model, @@ -7908,7 +8117,6 @@ 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) _complete_resolved_sampled_stops(tokenized_histories, builders) return ( TokenizedMultiHistoryTrajectory( diff --git a/src/art_inference/chat_template.py b/src/art_inference/chat_template.py index 93b18f294..adcfee50c 100644 --- a/src/art_inference/chat_template.py +++ b/src/art_inference/chat_template.py @@ -1,3 +1,4 @@ +from collections.abc import Sequence import json import re from typing import Any @@ -117,29 +118,159 @@ def operations(text: str): continue if not edits: return template - # Dropping structured reasoning must not trim the visible assistant body. + + def apply_edits() -> str: + result = template + for (start, end), replacement in sorted(edits.items(), reverse=True): + result = result[:start] + replacement + result[end:] + return result + + # Only the content assignment consumed by a recognized parser may lose + # its trim. Independent preview macros/branches have their own bindings. + try: + tree = WithoutWhitespace().visit(env.parse(template)) + except TemplateSyntaxError: + # Optional binding analysis must not undo independently proved parser + # edits when the renderer supports extensions absent from this parser. + return apply_edits() + assignments = list(tree.find_all(nodes.Assign)) + locations = [] + parsed_assignments = [] + for block in blocks: + start, body_start, body_end, end = block + if template[body_start:body_end].split(None, 1)[:1] != ["set"]: + continue + try: + body = env.parse(template[start:end]).body + except TemplateSyntaxError: + continue + if len(body) == 1 and isinstance(body[0], nodes.Assign): + locations.append(block) + parsed_assignments.append(body[0]) + if assignments != parsed_assignments: + return apply_edits() # Do not guess an assignment's source span. + # An AST-equivalent block with comments/data between its tags may not be + # one of the source spans removed above. Join If nodes in source order too. + conditions = [] + for block in blocks: + statement = template[block[1] : block[2]].strip() + keyword = statement.split(None, 1)[:1] + if keyword not in (["if"], ["elif"]): + continue + try: + parsed = env.parse( + "{% if " + statement.split(None, 1)[1] + " %}{% endif %}" + ) + except TemplateSyntaxError: + return apply_edits() + if len(parsed.body) != 1 or not isinstance(parsed.body[0], nodes.If): + return apply_edits() + conditions.append((block, parsed.body[0].test)) + branches = list(tree.find_all(nodes.If)) + if [node.test for node in branches] != [test for _, test in conditions]: + return apply_edits() + edited_parsers = { + id(node) + for node, (block, _) in zip(branches, conditions, strict=True) + if any(start == block[0] for start, _ in edits) + } + selected = set() + shared = set() + + def writes_content(node: nodes.Assign | nodes.AssignBlock) -> bool: + target = node.target + return ( + isinstance(target, nodes.Name) + and target.name == "content" + or any(n.name == "content" for n in target.find_all(nodes.Name)) + ) + + def reads_content(node: nodes.Node) -> bool: + return ( + isinstance(node, nodes.Name) + and node.name == "content" + and node.ctx == "load" + ) or any( + n.name == "content" and n.ctx == "load" for n in node.find_all(nodes.Name) + ) + + def assistant_condition(test: nodes.Node) -> bool | None: + # The rewrite changes assistant content only. Ignore paths proved to + # handle a different role, but inspect every unknown branch for users + # of the original trimmed value before the destructive parser. + if ( + isinstance(test, nodes.Compare) + and test.expr + == nodes.Getattr(nodes.Name("message", "load"), "role", "load") + and len(test.ops) == 1 + and test.ops[0].op in ("eq", "ne") + and isinstance(test.ops[0].expr, nodes.Const) + ): + equal = test.ops[0].expr.value == "assistant" + return equal if test.ops[0].op == "eq" else not equal + return None + + def visit(body: Sequence[nodes.Node], binding: nodes.Assign | None = None) -> None: + for node in body: + if isinstance(node, nodes.If): + if id(node) in edited_parsers and node == operation[0]: + if binding is not None: + selected.add(id(binding)) + binding = None + continue + if binding is not None and reads_content(node.test): + shared.add(id(binding)) + condition = assistant_condition(node.test) + if condition is not False: + visit(node.body, binding) + if condition is not True: + # elif_ is a list of If nodes whose else is stored on the + # outer If. Stop following it once a role match is proved. + for branch in node.elif_: + visit([branch], binding) + if assistant_condition(branch.test) is True: + break + else: + visit(node.else_, binding) + if any( + writes_content(n) + for n in node.find_all((nodes.Assign, nodes.AssignBlock)) + ): + binding = None + else: + if binding is not None and reads_content(node): + shared.add(id(binding)) + if isinstance(node, nodes.Assign): + if writes_content(node): + binding = node + else: + # Macro/loop/with/block bodies have independent bindings. + for _, value in node.iter_fields(): + if isinstance(value, list) and all( + isinstance(n, nodes.Node) for n in value + ): + visit(value) + if isinstance(node, nodes.AssignBlock) and writes_content(node): + binding = None + + visit(tree.body) trims = [ - env.parse("{% set content = " + content + " %}").body + env.parse("{% set content = " + content + " %}").body[0] 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 + for node, (start, body_start, body_end, end) in zip( + assignments, locations, strict=True + ): + if id(node) in selected - shared and node 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] + ) + return apply_edits() def chat_template_with_preserved_thinking(chat_template: object) -> object: @@ -151,7 +282,9 @@ def chat_template_with_preserved_thinking(chat_template: object) -> object: } if not isinstance(chat_template, str): return chat_template - chat_template = _without_inline_reasoning_parser(chat_template) + literal_template = _without_inline_reasoning_parser(chat_template) + inline_parser_removed = literal_template != chat_template + chat_template = literal_template replacements = ( ( _QWEN_DROP_PRIOR_THINKING, @@ -233,10 +366,13 @@ def chat_template_with_preserved_thinking(chat_template: object) -> object: "reasoning_content + '\\n\\n\\n'", "reasoning_content + ('\\n\\n' if preserve_thinking and message.reasoning_content is string and reasoning_content else '\\n\\n\\n')", ) - chat_template = chat_template.replace( - "set content = render_content(message.content, true)|trim", - "set content = (render_content(message.content, true) if preserve_thinking and message.role == 'assistant' else render_content(message.content, true)|trim)", - ) + if not inline_parser_removed: + # The recognized parser path already preserved its own input. Do + # not apply the legacy template-wide trim rewrite to other macros. + chat_template = chat_template.replace( + "set content = render_content(message.content, true)|trim", + "set content = (render_content(message.content, true) if preserve_thinking and message.role == 'assistant' else render_content(message.content, true)|trim)", + ) if "clear_thinking" in chat_template: chat_template = chat_template.replace( "{{ content.strip() }}", diff --git a/tests/unit/test_literal_reasoning_content.py b/tests/unit/test_literal_reasoning_content.py index 1c8056ec8..2a7f9bc30 100644 --- a/tests/unit/test_literal_reasoning_content.py +++ b/tests/unit/test_literal_reasoning_content.py @@ -419,3 +419,202 @@ def test_plain_block_whitespace_keeps_prior_inline_parser_coverage(separator, qu content = "HEADliteralTAIL" assert content in _render(fixed, [_USER, {"role": "assistant", "content": content}]) assert chat_template_with_preserved_thinking(fixed) == fixed + + +@pytest.mark.parametrize("preserve", [False, True]) +@pytest.mark.parametrize("newline", ["\n", "\r\n"]) +@pytest.mark.parametrize("layout", ["macros", "same_line", "branches"]) +def test_content_trim_is_scoped_to_the_recognized_parser(preserve, newline, layout): + match = _QWEN_INLINE_REASONING.search(_TEMPLATE) + assert match is not None + parser = match.group() + trim = "{% set content = render_content(message.content, true)|trim %}" + render = "{% macro render_content(content, count) %}{{ content }}{% endmacro %}" + preview = "{% macro preview(message) %}" + trim + "[{{ content }}]{% endmacro %}" + main = ( + "{% macro answer(message) %}" + trim + parser + "[{{ content }}]{% endmacro %}" + ) + if layout == "branches": + main = ( + "{% macro answer(message) %}{% if message.role == 'assistant' %}" + + trim + + parser + + "[{{ content }}]{% else %}" + + trim + + "[{{ content }}]{% endif %}{% endmacro %}" + ) + separator = "" if layout == "same_line" else newline + template = separator.join( + [render, preview, main, "{{ answer(message) }}|{{ preview(message) }}"] + ) + fixed = chat_template_with_preserved_thinking(template) + assert isinstance(fixed, str) + assert preview in fixed + env = ImmutableSandboxedEnvironment(trim_blocks=True, lstrip_blocks=True) + kwargs = dict( + message={"role": "assistant", "content": " answer "}, + preserve_thinking=preserve, + ) + assert env.from_string(fixed).render(**kwargs).endswith("[ answer ]|[answer]") + kwargs["message"]["content"] = " beforeliteralafter " + assert ( + env.from_string(fixed) + .render(**kwargs) + .endswith( + "[ beforeliteralafter ]|[beforeliteralafter]" + ) + ) + assert chat_template_with_preserved_thinking(fixed) == fixed + if layout == "branches": + kwargs["message"] = {"role": "user", "content": " answer "} + assert env.from_string(fixed).render(**kwargs).endswith("[answer]|[answer]") + + +@pytest.mark.parametrize("preserve", [False, True]) +@pytest.mark.parametrize("raw", [False, True]) +def test_qwen_content_preview_macro_is_unchanged(preserve, raw): + preview = "{% macro preview(message) %}{% set content = render_content(message.content, true)|trim %}[{{ content }}]{% endmacro %}" + body = _TEMPLATE + if raw: + body = body.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 %}", + ) + template = preview + body + "{{ preview(messages[-1]) }}" + fixed = chat_template_with_preserved_thinking(template) + assert isinstance(fixed, str) + assert preview in fixed + content = " beforeliteralafter " + rendered = _render( + fixed, + [_USER, {"role": "assistant", "content": content}], + preserve_thinking=preserve, + ) + assert content + "<|im_end|>\n" in rendered + assert rendered.endswith("[beforeliteralafter]") + + +@pytest.mark.parametrize("shadow", ["assignment", "conditional", "scope"]) +def test_content_trim_requires_a_proven_local_binding(shadow): + match = _QWEN_INLINE_REASONING.search(_TEMPLATE) + assert match is not None + parser = match.group() + trim = "{% set content = render_content(message.content, true)|trim %}" + render = "{% macro render_content(content, count) %}{{ content }}{% endmacro %}" + if shadow == "assignment": + template = render + trim + "{% set content = 'replacement' %}" + parser + elif shadow == "conditional": + template = ( + render + + trim + + "{% if custom %}{% set content = 'replacement' %}{% endif %}" + + parser + ) + else: + template = ( + render + + trim + + "{% macro nested() %}" + + parser + + "{{ content }}{% endmacro %}{{ nested() }}" + ) + fixed = _without_inline_reasoning_parser(template) + assert trim in fixed + assert not _QWEN_INLINE_REASONING.search(fixed) + assert _without_inline_reasoning_parser(fixed) == fixed + + +@pytest.mark.parametrize("extension", ["loopcontrols", "generation"]) +def test_optional_trim_binding_parse_keeps_proven_parser_removal(extension): + from jinja2 import nodes + from jinja2.ext import Extension + + class Generation(Extension): + tags = {"generation"} + + def parse(self, parser): + next(parser.stream) + body = parser.parse_statements(["name:endgeneration"], drop_needle=True) + return nodes.Scope(body) + + env = ImmutableSandboxedEnvironment( + extensions=["jinja2.ext.loopcontrols", Generation], + trim_blocks=True, + lstrip_blocks=True, + ) + match = _QWEN_INLINE_REASONING.search(_TEMPLATE) + assert match is not None + parser = match.group() + prefix = ( + "{% for item in [1] %}{% break %}{% endfor %}" + if extension == "loopcontrols" + else "{% generation %}PUBLIC{% endgeneration %}" + ) + template = prefix + "{% set content = message.content %}" + parser + "{{ content }}" + content = "HEADliteralTAIL" + kwargs = {"message": {"role": "assistant", "content": content}} + assert env.from_string(template).render(**kwargs).endswith("TAIL") + fixed = chat_template_with_preserved_thinking(template) + assert isinstance(fixed, str) + assert prefix in fixed + assert env.from_string(fixed).render(**kwargs).endswith(content) + assert chat_template_with_preserved_thinking(fixed) == fixed + + +@pytest.mark.parametrize("consumer", ["output", "alias", "condition", "other_branch"]) +def test_shared_content_preview_prevents_ambiguous_trim_rewrite(consumer): + match = _QWEN_INLINE_REASONING.search(_TEMPLATE) + assert match is not None + parser = match.group() + trim = "{% set content = render_content(message.content, true)|trim %}" + render = "{% macro render_content(content, count) %}{{ content }}{% endmacro %}" + if consumer == "output": + body = "[{{ content }}]" + parser + "[{{ content }}]" + elif consumer == "alias": + body = "{% set preview = content %}" + parser + "[{{ preview }}][{{ content }}]" + elif consumer == "condition": + body = "{% if content %}PREVIEW{% endif %}" + parser + "[{{ content }}]" + else: + body = ( + "{% if preview_only %}[{{ content }}]{% else %}" + + parser + + "[{{ content }}]{% endif %}" + ) + template = render + trim + body + fixed = chat_template_with_preserved_thinking(template) + assert isinstance(fixed, str) + assert trim in fixed + assert not _QWEN_INLINE_REASONING.search(fixed) + env = ImmutableSandboxedEnvironment(trim_blocks=True, lstrip_blocks=True) + kwargs = dict(message={"role": "assistant", "content": " "}, preview_only=True) + assert env.from_string(fixed).render(**kwargs) == env.from_string(template).render( + **kwargs + ) + + +def test_only_source_edited_parser_authorizes_its_content_trim(): + match = _QWEN_INLINE_REASONING.search(_TEMPLATE) + assert match is not None + parser = match.group() + unedited = parser.replace("%}", "%}{# preserve this custom block #}", 1) + trim = "{% set content = render_content(message.content, true)|trim %}" + render = "{% macro render_content(content, count) %}{{ content }}{% endmacro %}" + preview = ( + "{% macro preview(message) %}" + + trim + + unedited + + "[{{ content }}]{% endmacro %}" + ) + answer = ( + "{% macro answer(message) %}" + trim + parser + "[{{ content }}]{% endmacro %}" + ) + template = ( + render + preview + answer + "{{ answer(message) }}|{{ preview(message) }}" + ) + fixed = _without_inline_reasoning_parser(template) + assert preview in fixed + env = ImmutableSandboxedEnvironment(trim_blocks=True, lstrip_blocks=True) + rendered = env.from_string(fixed).render( + message={"role": "assistant", "content": " HEADliteralTAIL "} + ) + assert rendered == "[ HEADliteralTAIL ]|[TAIL]" diff --git a/tests/unit/trajectories/test_evidence_reuse.py b/tests/unit/trajectories/test_evidence_reuse.py index 31ff8d55c..c8cee9a0c 100644 --- a/tests/unit/trajectories/test_evidence_reuse.py +++ b/tests/unit/trajectories/test_evidence_reuse.py @@ -509,3 +509,93 @@ def test_edited_history_does_not_reuse_projection_proof(edit): history.messages[0]["content"] = "new question" with pytest.raises((ValueError, AssertionError)): history.tokenize(tokenizer=_CharacterTemplateTokenizer()) + + +@pytest.mark.parametrize("mutation", ["model", "logprob"]) +def test_boundary_decoder_cannot_change_later_source_model(monkeypatch, mutation): + from test_tokenize import _character_template_history + + history, tokenizer, _ = _character_template_history() + sources = [ + source + for source in history.message_sources + if source is not None and source.choice_index is not None + ] + middle = sources[1].exchange + final = sources[2].exchange + assert isinstance(middle, tr.ChatCompletionsExchange) + output = extras(middle)["token_ids"] + original_decode = tokenizer.decode + calls = [] + + def decode(ids, **kwargs): + if ids == output: + calls.append(True) + if mutation == "model": + final.request["model"] = "changed/model" + else: + assert isinstance(final, tr.ChatCompletionsExchange) + logprobs(final)[0].logprob = -99 + return original_decode(ids, **kwargs) + + monkeypatch.setattr(tokenizer, "decode", decode) + with pytest.raises( + ValueError, + match="context changed|model no longer matches|Sampled source changed", + ): + history.tokenize(tokenizer=tokenizer) + assert calls + + +def test_stop_callback_cannot_change_next_request_then_restore_it(): + 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": + second.request["tools"] = [ + {"type": "function", "function": {"name": "changed"}} + ] + else: + second.request.pop("tools", None) + return {"input_ids": [3 if text == "first-stop" else 6]} + + with pytest.raises( + ValueError, match="context changed during tokenization callback" + ): + trajectory(first, second).tokenize(tokenizer=StopTokenizer()) + assert calls == ["first-stop"] + + +@pytest.mark.parametrize("kind", [RuntimeError, KeyboardInterrupt, SystemExit]) +def test_boundary_decoder_mutation_keeps_uncaught_exception_identity(monkeypatch, kind): + from test_tokenize import _character_template_history + + history, tokenizer, _ = _character_template_history() + selected = [ + source + for source in history.message_sources + if source is not None and source.choice_index is not None + ] + middle, final = selected[1].exchange, selected[2].exchange + assert isinstance(middle, tr.ChatCompletionsExchange) + assert isinstance(final, tr.ChatCompletionsExchange) + output = extras(middle)["token_ids"] + error = kind("public callback failure") + original_decode = tokenizer.decode + + def decode(ids, **kwargs): + if ids == output: + logprobs(final)[0].logprob = -99 + raise error + return original_decode(ids, **kwargs) + + monkeypatch.setattr(tokenizer, "decode", decode) + with pytest.raises(kind) as caught: + history.tokenize(tokenizer=tokenizer) + assert caught.value is error diff --git a/tests/unit/trajectories/test_recorded_prompt_roles.py b/tests/unit/trajectories/test_recorded_prompt_roles.py index 85bdcd387..12e181255 100644 --- a/tests/unit/trajectories/test_recorded_prompt_roles.py +++ b/tests/unit/trajectories/test_recorded_prompt_roles.py @@ -430,3 +430,95 @@ def malformed(self, text, **kwargs): ) is None ) + + +@pytest.mark.parametrize("offsets", [False, True]) +def test_interior_historical_parser_cannot_change_sample_conditioning( + monkeypatch, offsets +): + from openai.types.chat import ChatCompletion + from test_literal_thinking_off import _TemplateTokenizer + from test_tokenize import _chat_exchange + + from art_inference.chat_template import _QWEN_INLINE_STATEMENTS + + template = ( + "{% for message in messages %}{% set content = message.content %}" + + "".join("{% " + part + " %}" for part in _QWEN_INLINE_STATEMENTS) + + "{{ content }}{% if message.role == 'assistant' %}§{% endif %}{% endfor %}" + ) + + class Tokenizer(_TemplateTokenizer): + eos_token_id = ord("§") + + def __call__(self, text, **kwargs): + if kwargs.get("return_offsets_mapping") and not offsets: + raise NotImplementedError("No public offsets") + return super().__call__(text, **kwargs) + + tokenizer = Tokenizer() + tokenizer.chat_template = template + messages = [{"role": "user", "content": "first query"}] + exchanges = [] + records = [] + for index in range(2): + prompt = tokenizer.apply_chat_template(messages, add_generation_prompt=True) + assert isinstance(prompt, list) + output = list(map(ord, f"answer{index}§")) + exchange = _chat_exchange(prompt, output, offset=index) + exchange.request["messages"] = cast( + list[ChatCompletionMessageParam], deepcopy(messages) + ) + exchange.request["chat_template"] = template + payload = exchange.response.model_dump(mode="python") + payload["choices"][0]["message"]["content"] = f"answer{index}" + exchange.response = ChatCompletion.model_validate(payload) + exchanges.append(exchange) + records.append((prompt, output)) + messages.extend( + [ + {"role": "assistant", "content": f"answer{index}"}, + {"role": "user", "content": "middle query"}, + {"role": "assistant", "content": "xy"}, + {"role": "user", "content": "next query"}, + ] + ) + value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(chat_completions=exchanges)) + original = value.model_dump() + if not offsets: + with pytest.raises(ValueError, match="Cannot preserve request roles"): + value.tokenize(tokenizer=tokenizer, multi_history=True) + assert value.model_dump() == original + return + actual = value.tokenize(tokenizer=tokenizer, multi_history=True) + assert len(actual.histories) == 1 + result = actual.histories[0] + assert result.tokens == records[-1][0] + records[-1][1] + for prompt, output in records: + assert result.tokens[: len(prompt)] == prompt + assert result.tokens[len(prompt) : len(prompt) + len(output)] == output + assert all( + flag & tr.TokenFlag.SAMPLED + for flag in result.flags[len(prompt) : len(prompt) + len(output)] + ) + assert value.model_dump() == original + + with monkeypatch.context() as patch: + patch.setattr( + module, "chat_template_with_preserved_thinking", lambda value: value + ) + expected = value.tokenize(tokenizer=tokenizer, multi_history=True) + assert result.flags == expected.histories[0].flags + assert all( + a == b or math.isnan(a) and math.isnan(b) + for a, b in zip(result.logprobs, expected.histories[0].logprobs, strict=True) + ) + for flag in ( + tr.TokenFlag.SAMPLED, + tr.TokenFlag.OUTPUT, + tr.TokenFlag.ASSISTANT, + tr.TokenFlag.STOP, + ): + assert tr.first_occurrence_masks( + actual.histories, where=flag + ) == tr.first_occurrence_masks(expected.histories, where=flag) diff --git a/tests/unit/trajectories/test_resolved_stop_authority.py b/tests/unit/trajectories/test_resolved_stop_authority.py index fdfd2cf90..b32f6be19 100644 --- a/tests/unit/trajectories/test_resolved_stop_authority.py +++ b/tests/unit/trajectories/test_resolved_stop_authority.py @@ -131,6 +131,57 @@ def apply_chat_template(self, messages, **kwargs): ) +@pytest.mark.parametrize("supplied", [False, True]) +@pytest.mark.parametrize( + "mutation", + ["request_role", "request_tools", "history_role", "history_kwargs", "kwargs_order"], +) +def test_later_history_callback_cannot_change_completed_context( + monkeypatch, supplied, mutation +): + first = branch(0, length=False, model="model/a") + second = branch(1, length=True, model="model/b") + + value = trajectory(first, second) + histories = value.histories() + original_histories = type(value).histories + monkeypatch.setattr( + type(value), + "histories", + lambda self, **kwargs: ( + histories if self is value else original_histories(self, **kwargs) + ), + ) + options = {"first": 1, "second": 2} + + class Tokenizer(_CharacterTemplateTokenizer): + def apply_chat_template(self, messages, **kwargs): + if mutation == "request_role": + first.request["messages"][0]["role"] = "assistant" + elif mutation == "request_tools": + first.request["tools"] = [ + {"type": "function", "function": {"name": "changed"}} + ] + elif mutation == "history_role": + histories[0].messages[0]["role"] = "assistant" + elif mutation == "history_kwargs": + histories[0].chat_template_kwargs = {"enable_thinking": True} + else: + options["first"] = options.pop("first") + return super().apply_chat_template(messages, **kwargs) + + tokenizer = Tokenizer() + bind(monkeypatch, {"model/a": tokenizer, "model/b": tokenizer}) + with pytest.raises( + ValueError, match="context changed during tokenization callback" + ): + value.tokenize( + multi_history=True, + tokenizer=tokenizer if supplied else None, + chat_template_kwargs=options if mutation == "kwargs_order" else None, + ) + + def test_stop_postpass_checks_original_evidence_before_consuming_it(monkeypatch): first = branch(0, length=False) second = branch(1, length=False) @@ -398,3 +449,138 @@ async def test_public_async_default_dispatch_uses_resolved_stop_authority(monkey assert loads == ["test/model"] assert len(results) == 1 and results[0].trajectory is value assert results[0].histories[-1].flags[-1] & tr.TokenFlag.STOP + + +@pytest.mark.parametrize("mapping_type", ["ordered", "proxy"]) +@pytest.mark.parametrize("mutation", [None, "order", "value"]) +def test_standard_mapping_options_keep_context_proof(mapping_type, mutation): + from collections import OrderedDict + from types import MappingProxyType + + options = OrderedDict(first=1, second=2) + selected = options if mapping_type == "ordered" else MappingProxyType(options) + exchange = branch(0, length=False) + + class Tokenizer(_CharacterTemplateTokenizer): + changed = False + + def apply_chat_template(self, messages, **kwargs): + if not self.changed and mutation is not None: + self.changed = True + if mutation == "order": + options.move_to_end("first") + else: + options["first"] = 99 + return super().apply_chat_template(messages, **kwargs) + + value = trajectory(exchange) + tokenizer = Tokenizer() + if mutation is None: + actual = value.tokenize(tokenizer=tokenizer, chat_template_kwargs=selected) + expected = value.tokenize( + tokenizer=_CharacterTemplateTokenizer(), chat_template_kwargs=dict(selected) + ) + assert actual.tokens == expected.tokens and actual.flags == expected.flags + assert all( + a == b or math.isnan(a) and math.isnan(b) + for a, b in zip(actual.logprobs, expected.logprobs, strict=True) + ) + else: + with pytest.raises( + ValueError, match="context changed during tokenization callback" + ): + value.tokenize(tokenizer=tokenizer, chat_template_kwargs=selected) + assert tokenizer.changed + + +@pytest.mark.parametrize( + "protocol", ["messages", "responses", "completion_tokens", "completion_string"] +) +@pytest.mark.parametrize("mutate", [False, True]) +def test_protocol_history_context_survives_stop_callbacks(protocol, mutate): + from test_tokenize import ( + _completion_exchange, + _message_exchange, + _response_exchange, + ) + + if protocol == "messages": + exchange = _message_exchange( + tr.MessagesRequest( + model="test/model", messages=[{"role": "user", "content": "question"}] + ), + prompt_token_ids=[1], + token_ids=[2], + logprobs=[-0.2], + ) + exchanges = tr.TrajectoryExchanges(messages=[exchange]) + elif protocol == "responses": + exchange = _response_exchange("public-response", 2, prompt_token_ids=[1]) + exchanges = tr.TrajectoryExchanges(responses=[exchange]) + else: + exchange = _completion_exchange( + prompt=[1] if protocol == "completion_tokens" else "question" + ) + exchanges = tr.TrajectoryExchanges(completions=[exchange]) + history = tr.Trajectory(exchanges=exchanges).histories()[0] + assert not isinstance(history, tr.LegacyHistory) + calls = [] + + class Tokenizer: + eos_token_id = 2 + + def __call__(self, text, **kwargs): + raise AssertionError("Complete native records need no encoding") + + def convert_tokens_to_ids(self, token): + calls.append(token) + if mutate: + if protocol == "messages": + assert isinstance(history, tr.AnthropicMessagesHistory) + history.system = "changed system" + elif protocol == "responses": + assert isinstance(history, tr.ResponsesHistory) + history.previous_response_id = "changed-context" + else: + assert isinstance( + history, + (tr.CompletionsTokenHistory, tr.CompletionsStringHistory), + ) + history.sampled_spans = [] + return None + + if mutate: + with pytest.raises( + ValueError, match="context changed during tokenization callback" + ): + history.tokenize(tokenizer=cast(Any, Tokenizer())) + else: + actual = history.tokenize(tokenizer=cast(Any, Tokenizer())) + assert actual.tokens == [1, 2] + assert actual.flags[-1] & tr.TokenFlag.SAMPLED + assert calls + + +def test_context_snapshot_shared_values_are_fresh_between_observations(): + from copy import deepcopy + + shared = {"role": "assistant", "content": ["original"]} + context = [[shared], {"messages": [shared]}] + before = module._tokenization_context(context) + assert before == module._tokenization_context(deepcopy(context)) + validate = module._tokenization_context_validator(context) + shared["content"][0] = "changed" + assert module._tokenization_context(context) != before + with pytest.raises(ValueError, match="context changed"): + validate(True) + shared["content"][0] = "original" + validate(True) + + +def test_context_snapshot_does_not_certify_cyclic_or_opaque_values(): + cyclic = [] + cyclic.append(cyclic) + with pytest.raises(RecursionError): + module._tokenization_context(cyclic) + with pytest.raises(TypeError, match="Unsupported mutable"): + module._tokenization_context([object()]) From bd64400f08502c5b2dfc118358abd2a9b0a53a34 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 20:00:31 +0000 Subject: [PATCH 18/67] Validate all consumed sources after historical role proofs --- src/art/trajectories/_tokenize.py | 21 ++- .../test_recorded_prompt_roles.py | 124 ++++++++++++++++-- 2 files changed, 134 insertions(+), 11 deletions(-) diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index 0d9c0d4ba..aa9122ca3 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -626,9 +626,21 @@ def check_context() -> None: "Renderer changed context while proving recorded request roles" ) + normalization_template = template + if isinstance(getattr(tokenizer, "chat_template", None), dict) and callable( + select := getattr(tokenizer, "get_chat_template", None) + ): + normalization_template = select( + chat_template=template if isinstance(template, str) else None, + tools=tools, + ) + check_context() + 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), + normalize_tool_call_arguments_for_chat_template( + selected, normalization_template + ), tools=tools, tokenize=False, add_generation_prompt=add_generation_prompt, @@ -5379,6 +5391,10 @@ def _tokenize_chat_view( _trace: _TraceBuilder | None = None, _prior: Sequence[tuple[TokenizedHistory, _HistoryTokenizationTrace]] = (), ) -> TokenizedHistory: + if _trace is not None: + # Rendering may continue after exact assembly to prove historical roles. + # Keep every consumed source bound even for a standalone/single history. + _trace.track_sources = True _validate_history_sources(history) config = ( _TokenizerConfig(base_model or history.model or "") @@ -7909,8 +7925,7 @@ def tokenize_history( _projection_validated=_projection_validated, _copied_context=bool(copied), ) - if trace_builder.tokenizer is not None: - trace_builder.validate_context(True) + _validate_completed_sources([trace_builder]) if copied: if trace_builder is None or trace_builder.trace is None: raise ValueError( diff --git a/tests/unit/trajectories/test_recorded_prompt_roles.py b/tests/unit/trajectories/test_recorded_prompt_roles.py index 12e181255..d425f042b 100644 --- a/tests/unit/trajectories/test_recorded_prompt_roles.py +++ b/tests/unit/trajectories/test_recorded_prompt_roles.py @@ -432,10 +432,7 @@ def malformed(self, text, **kwargs): ) -@pytest.mark.parametrize("offsets", [False, True]) -def test_interior_historical_parser_cannot_change_sample_conditioning( - monkeypatch, offsets -): +def _interior_historical_case(offsets=True, selection="literal"): from openai.types.chat import ChatCompletion from test_literal_thinking_off import _TemplateTokenizer from test_tokenize import _chat_exchange @@ -448,42 +445,90 @@ def test_interior_historical_parser_cannot_change_sample_conditioning( + "{{ content }}{% if message.role == 'assistant' %}§{% endif %}{% endfor %}" ) + if selection != "literal": + template = template.replace( + "{{ content }}", + "{{ content }}{% for tool_call in message.tool_calls or [] %}{% for key, arg in tool_call.function.arguments.items() %}{{ key }}={{ arg }}{% endfor %}{% endfor %}", + ) + class Tokenizer(_TemplateTokenizer): + chat_template: Any eos_token_id = ord("§") + def get_chat_template(self, chat_template=None, tools=None): + if isinstance(self.chat_template, dict): + return self.chat_template.get(chat_template or "default", chat_template) + return chat_template or self.chat_template + + def apply_chat_template(self, messages, **kwargs): + selected = self.get_chat_template( + kwargs.pop("chat_template", None), kwargs.get("tools") + ) + return super().apply_chat_template( + messages, chat_template=selected, **kwargs + ) + def __call__(self, text, **kwargs): if kwargs.get("return_offsets_mapping") and not offsets: raise NotImplementedError("No public offsets") return super().__call__(text, **kwargs) tokenizer = Tokenizer() - tokenizer.chat_template = template + tokenizer.chat_template = ( + template if selection == "literal" else {"default": template, "named": template} + ) messages = [{"role": "user", "content": "first query"}] exchanges = [] records = [] for index in range(2): - prompt = tokenizer.apply_chat_template(messages, add_generation_prompt=True) + prompt = tokenizer.apply_chat_template( + module.normalize_tool_call_arguments_for_chat_template(messages, template), + add_generation_prompt=True, + ) assert isinstance(prompt, list) output = list(map(ord, f"answer{index}§")) exchange = _chat_exchange(prompt, output, offset=index) exchange.request["messages"] = cast( list[ChatCompletionMessageParam], deepcopy(messages) ) - exchange.request["chat_template"] = template + if selection is not None: + exchange.request["chat_template"] = ( + template if selection == "literal" else selection + ) payload = exchange.response.model_dump(mode="python") payload["choices"][0]["message"]["content"] = f"answer{index}" exchange.response = ChatCompletion.model_validate(payload) exchanges.append(exchange) records.append((prompt, output)) + historical: dict[str, Any] = {"role": "assistant", "content": "xy"} + if selection != "literal": + historical["tool_calls"] = [ + { + "id": "public-tool", + "type": "function", + "function": { + "name": "public", + "arguments": '{"public": "argument"}', + }, + } + ] messages.extend( [ {"role": "assistant", "content": f"answer{index}"}, {"role": "user", "content": "middle query"}, - {"role": "assistant", "content": "xy"}, + historical, {"role": "user", "content": "next query"}, ] ) value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(chat_completions=exchanges)) + return value, tokenizer, records + + +@pytest.mark.parametrize("offsets", [False, True]) +def test_interior_historical_parser_cannot_change_sample_conditioning( + monkeypatch, offsets +): + value, tokenizer, records = _interior_historical_case(offsets) original = value.model_dump() if not offsets: with pytest.raises(ValueError, match="Cannot preserve request roles"): @@ -522,3 +567,66 @@ def __call__(self, text, **kwargs): assert tr.first_occurrence_masks( actual.histories, where=flag ) == tr.first_occurrence_masks(expected.histories, where=flag) + + +@pytest.mark.parametrize("route", ["history", "single", "multi"]) +def test_historical_role_callback_cannot_replace_consumed_earlier_logprob( + monkeypatch, route +): + value, tokenizer, _ = _interior_historical_case() + first = value.exchanges.chat_completions[0] + apply = tokenizer.apply_chat_template + calls = [] + + def mutate(messages, **kwargs): + if kwargs.get("chat_template") == tokenizer.chat_template and len(messages) > 3: + calls.append(True) + logprobs = first.response.choices[0].logprobs + assert logprobs is not None and logprobs.content is not None + logprobs.content[0].logprob -= 0.5 + return apply(messages, **kwargs) + + monkeypatch.setattr(tokenizer, "apply_chat_template", mutate) + with pytest.raises(ValueError, match="[Ss]ampled source changed"): + if route == "history": + value.chat_completions_history().tokenize(tokenizer=tokenizer) + else: + value.tokenize(tokenizer=tokenizer, multi_history=route == "multi") + assert calls + + +@pytest.mark.parametrize("selection", [None, "named"]) +def test_named_original_request_proof_normalizes_tool_arguments(selection): + value, tokenizer, records = _interior_historical_case(selection=selection) + before = value.model_dump() + result = value.tokenize(tokenizer=tokenizer, multi_history=True) + assert result.histories[0].tokens == records[-1][0] + records[-1][1] + assert value.model_dump() == before + literal = value.model_copy(deep=True) + for exchange in literal.exchanges.chat_completions: + exchange.request["chat_template"] = tokenizer.chat_template["default"] + expected = literal.tokenize(tokenizer=tokenizer, multi_history=True) + assert result.histories[0].flags == expected.histories[0].flags + assert all( + a == b or math.isnan(a) and math.isnan(b) + for a, b in zip( + result.histories[0].logprobs, expected.histories[0].logprobs, strict=True + ) + ) + for flag in ( + tr.TokenFlag.SAMPLED, + tr.TokenFlag.OUTPUT, + tr.TokenFlag.ASSISTANT, + tr.TokenFlag.STOP, + ): + assert tr.first_occurrence_masks( + result.histories, where=flag + ) == tr.first_occurrence_masks(expected.histories, where=flag) + for prompt, output in records: + history = result.histories[0] + assert history.tokens[: len(prompt)] == prompt + assert history.tokens[len(prompt) : len(prompt) + len(output)] == output + assert all( + flag & tr.TokenFlag.SAMPLED + for flag in history.flags[len(prompt) : len(prompt) + len(output)] + ) From 5f39e1e927a757bdedaab78bcc16907007a1b0c4 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 20:27:23 +0000 Subject: [PATCH 19/67] Validate consumed evidence through rendered token replacement --- docs/features/additional-histories.mdx | 3 +- src/art/trajectories/_tokenize.py | 134 +++++++++++++----- .../unit/trajectories/test_evidence_reuse.py | 95 +++++++++++++ .../unit/trajectories/test_native_terminal.py | 23 +++ .../test_resolved_stop_authority.py | 47 ++++++ 5 files changed, 263 insertions(+), 39 deletions(-) diff --git a/docs/features/additional-histories.mdx b/docs/features/additional-histories.mdx index f15043116..11bacd3d2 100644 --- a/docs/features/additional-histories.mdx +++ b/docs/features/additional-histories.mdx @@ -189,7 +189,8 @@ 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. With a supplied renderer, request-owned assistant roles are proved throughout -the native stream, including between sampled responses. A request whose text +the native stream, including between sampled responses. This requirement applies +even when template normalization makes no change. 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. When literal-content normalization changes historical rendering, ART requires the original request diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index aa9122ca3..a11b79f33 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -6,6 +6,7 @@ from copy import deepcopy from dataclasses import dataclass, replace from datetime import datetime +from enum import Enum from functools import lru_cache from hashlib import sha256 import json @@ -968,6 +969,8 @@ def set( *, tokenizer: Tokenizer | None = None, ) -> None: + if self.validate_sources is not None: + self.validate_sources(None) self.tokenizer = tokenizer trace = _HistoryTokenizationTrace(source_keys=source_keys, sources=sources) trace.validate(tokenized) @@ -3463,6 +3466,28 @@ def _history_has_length_stop( return False +def _native_nonterminal_stops_known( + history: ChatCompletionsHistory, + *, + _fingerprints: dict[tuple[int, str, int], tuple[Exchange, str]] | None = None, +) -> bool: + sources: dict[_SampledSourceKey, object] = {} + for message, source in zip(history.messages, history.message_sources, strict=True): + if message.get("role") != "assistant": + continue + if source is None or not _source_is_sampled(source): + return False + sources[_sampled_source_key(source, _fingerprints=_fingerprints)] = source + for key, source in list(sources.items())[:-1]: + kind, reason = _source_stop_evidence(source, key) + if kind != "stop" or not isinstance(reason, int) or isinstance(reason, bool): + return False + output = _source_output_tokens(source, key) + if not output or output[-1] != reason: + return False + return bool(sources) + + def _history_needs_synthetic_stop( history: History, tokenizer: Tokenizer | None ) -> bool: @@ -4326,6 +4351,26 @@ def snapshot(item: object) -> object: return kind, item if kind is float: return kind, repr(item) + if isinstance(item, Enum): + if getattr(item, "__objclass__", kind) is not kind: + raise TypeError("Unsupported enum tokenization context") + return kind, snapshot( + { + key: child + for key, child in vars(item).items() + if key != "__objclass__" + } + ) + if isinstance(item, (str, int, float, bytes)): + if isinstance(item, str): + scalar = str.__str__(item) + elif isinstance(item, int): + scalar = int.__int__(item) + elif isinstance(item, float): + scalar = repr(float.__float__(item)) + else: + scalar = bytes.__bytes__(item) + return kind, scalar, snapshot(getattr(item, "__dict__", None)) previous = observed.get(id(item)) if previous is not None and previous[0] is item: return previous[1] @@ -4366,7 +4411,9 @@ def _tokenization_context_validator(value: object) -> Callable[[bool], None]: def validate(require_supported: bool) -> None: if expected is None and not require_supported: return - if expected is None or _tokenization_context(value) != expected: + if expected is None: + raise ValueError("Tokenization context cannot be checked for callbacks") + if _tokenization_context(value) != expected: raise ValueError( "Tokenization context changed during tokenization callback" ) @@ -5490,6 +5537,37 @@ def render_text( ): return recorded + # Generic rendering consumes the complete projected response set. Retain + # that evidence before callbacks, rather than blessing fresh keys afterward. + consumed_keys = { + id(source): _sampled_source_key(source) + for source in history.message_sources + if source is not None and _source_is_sampled(source) + } + consumed_sources = { + consumed_keys[id(source)]: source + for source in history.message_sources + if id(source) in consumed_keys + } + validate_consumed = _sampled_source_validator(consumed_sources) + if _trace is not None: + if _trace.validate_sources is not None: + _trace.validate_sources(None) + _trace.validate_sources = validate_consumed + + def consumed_source_key(source: object) -> _SampledSourceKey: + key = consumed_keys[id(source)] + validate_consumed(key) + return key + + def sampled_stop_suffix(tokens: Sequence[int], source: object) -> int: + key = consumed_source_key(source) + count = _sampled_stop_suffix( + tokens, source=source, source_key=key, tokenizer=resolved_tokenizer + ) + validate_consumed(key) + return count + prefix_render_cache = _PrefixChatRenderCache(render_normalized_text) def segmented_render( @@ -5708,6 +5786,8 @@ def part_ids(text: str) -> list[int]: output_cache: dict[int, tuple[list[int] | None, list[float]]] = {} def source_prompt_tokens(source: object) -> list[int] | None: + if id(source) in consumed_keys: + consumed_source_key(source) key = id(source) if key not in prompt_cache: prompt_cache[key] = _chat_source_prompt_tokens(source) @@ -5716,6 +5796,8 @@ def source_prompt_tokens(source: object) -> list[int] | None: def source_output_tokens( source: object, ) -> tuple[list[int] | None, list[float]]: + if id(source) in consumed_keys: + consumed_source_key(source) key = id(source) if key not in output_cache: output_cache[key] = _chat_source_full_tokens(source) @@ -5842,6 +5924,7 @@ def source_matches_context(source: object) -> bool: ) break + validate_consumed(None) canonical_length_stop_mask = _synthetic_length_stop_mask( messages, history.message_sources, @@ -6377,7 +6460,7 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: for position, message_index in enumerate(sampled_message_indices): source = history.message_sources[message_index] assert source is not None - source_key = _sampled_source_key(source) + source_key = consumed_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" @@ -6387,12 +6470,7 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: stop_reason == "stop" and bool(_terminator_ids(resolved_tokenizer)) and output is not None - and not _sampled_stop_suffix( - output, - source=source, - source_key=source_key, - tokenizer=resolved_tokenizer, - ) + and not sampled_stop_suffix(output, source) ) if synthetic_stop and position + 1 < len(sampled_message_indices): length_stop_boundaries_complete = False @@ -6618,7 +6696,7 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: and chat_template is None and chat_template_kwargs is None and source_matches_context(source) - and _source_stop_evidence(source, _sampled_source_key(source))[0] + and _source_stop_evidence(source, consumed_source_key(source))[0] != "length" ): exact_output_matches = locations(full_exact, search_cursor) @@ -6797,7 +6875,7 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: if len(full_logprobs) == len(full_exact) else [math.nan] * len(full_exact), True, - _sampled_source_key(source), + consumed_source_key(source), source, None, ) @@ -6831,7 +6909,7 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: rendered[start:end], [math.nan] * (end - start), False, - _sampled_source_key(source), + consumed_source_key(source), source, None, ) @@ -6882,12 +6960,7 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: multi_generation_response or len(parts) != 1 or parts[0][0] != "content" ): start = generation_start - if _sampled_stop_suffix( - full_exact, - source=source, - source_key=_sampled_source_key(source), - tokenizer=resolved_tokenizer, - ): + if sampled_stop_suffix(full_exact, source): # Adjacent assistants can share a role mask. Prove this message's # end before replacing its rendered closing markup and stop. completed = probe_render( @@ -6921,7 +6994,7 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: if len(full_logprobs) == len(full_exact) else [math.nan] * len(full_exact), True, - _sampled_source_key(source), + consumed_source_key(source), source, None, ) @@ -7005,12 +7078,7 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: elif ( exact is not None and corrected_message_end is not None - and _sampled_stop_suffix( - exact, - source=source, - source_key=_sampled_source_key(source), - tokenizer=resolved_tokenizer, - ) + and sampled_stop_suffix(exact, source) ): # Use the proven message end to replace its rendered stop, # just as the whole-message path does for sampled stops. @@ -7028,7 +7096,7 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: if ( exact is not None and corrected_message_end is not None - and _source_stop_evidence(source, _sampled_source_key(source))[0] + and _source_stop_evidence(source, consumed_source_key(source))[0] != "length" ): # Source evidence assigns STOP; retain synthetic length boundaries. @@ -7068,7 +7136,7 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: if len(logprobs) == len(replacement) else [math.nan] * len(replacement), exact is not None, - _sampled_source_key(source), + consumed_source_key(source), source, span[1] if corrected_message_end is not None @@ -7267,6 +7335,7 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: flags[index] |= TokenFlag.EXACT if history.model is None: raise ValueError("History tokenization requires a model") + validate_consumed(None) _mark_sampled_stops( token_ids, flags, @@ -7735,18 +7804,7 @@ def _tokenize_history( ( 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 _native_nonterminal_stops_known(history, _fingerprints=fingerprints) ) and not needs_request_roles and not needs_synthetic_stop diff --git a/tests/unit/trajectories/test_evidence_reuse.py b/tests/unit/trajectories/test_evidence_reuse.py index c8cee9a0c..c30e2d597 100644 --- a/tests/unit/trajectories/test_evidence_reuse.py +++ b/tests/unit/trajectories/test_evidence_reuse.py @@ -599,3 +599,98 @@ def decode(ids, **kwargs): with pytest.raises(kind) as caught: history.tokenize(tokenizer=tokenizer) assert caught.value is error + + +@pytest.mark.parametrize("mutate", [False, True]) +def test_explicit_render_binds_logprobs_before_stop_encoder(mutate): + from test_tokenize import _CharacterTemplateTokenizer + + calls = [] + exchange = None + + class Tokenizer(_CharacterTemplateTokenizer): + eos_token_id = None + all_special_tokens = [] + + def convert_tokens_to_ids(self, token): + return None + + def __call__(self, text, **kwargs): + if text == "END": + calls.append(text) + if mutate and len(calls) == 1: + assert exchange is not None + logprobs(exchange)[0].logprob = -9.0 + return {"input_ids": [99]} + return super().__call__(text, **kwargs) + + tokenizer = Tokenizer() + messages = [{"role": "user", "content": "turn 0"}] + prompt = tokenizer.apply_chat_template(messages, add_generation_prompt=True) + assert isinstance(prompt, list) + exchange = _chat_exchange(prompt, [20]) + choice = exchange.response.choices[0] + choice.message.content = "ab" + logprobs(exchange)[0].logprob = -0.5 + extras(exchange)["stop_reason"] = "END" + history = tr.Trajectory( + exchanges=tr.TrajectoryExchanges(chat_completions=[exchange]) + ).chat_completions_history() + if mutate: + with pytest.raises(ValueError, match="[Ss]ampled source changed"): + history.tokenize(tokenizer=tokenizer, chat_template="explicit override") + else: + result = history.tokenize( + tokenizer=tokenizer, chat_template="explicit override" + ) + assert result.tokens[len(prompt)] == 20 and result.logprobs[len(prompt)] == -0.5 + assert result.flags[len(prompt)] & tr.TokenFlag.SAMPLED + assert calls + + +@pytest.mark.parametrize("mutate", [False, True]) +def test_explicit_render_checks_next_source_before_stop_consumption(mutate): + from test_tokenize import _CharacterTemplateTokenizer + + calls = [] + second = None + + class Tokenizer(_CharacterTemplateTokenizer): + eos_token_id = None + all_special_tokens = [] + + def convert_tokens_to_ids(self, token): + return None + + def __call__(self, text, **kwargs): + if text in ("FIRST", "SECOND", "CHANGED"): + calls.append(text) + assert second is not None + if text == "FIRST" and mutate: + extras(second)["stop_reason"] = "CHANGED" + elif text == "CHANGED": + # A final-only equality check could miss this restoration. + extras(second)["stop_reason"] = "SECOND" + return {"input_ids": [99]} + return super().__call__(text, **kwargs) + + tokenizer = Tokenizer() + prompt = tokenizer._encode("turn 0") + first = _chat_exchange(prompt, [20]) + second = _chat_exchange([*prompt, 20, *tokenizer._encode("turn 1")], [21], offset=1) + extras(first)["stop_reason"] = "FIRST" + extras(second)["stop_reason"] = "SECOND" + value = trajectory(first, second) + before = value.model_dump_json() + if mutate: + with pytest.raises(ValueError, match="[Ss]ampled source changed"): + value.tokenize(tokenizer=tokenizer, chat_template="explicit override") + assert "FIRST" in calls and "CHANGED" not in calls + else: + result = value.tokenize(tokenizer=tokenizer, chat_template="explicit override") + assert [ + token + for token, flag in zip(result.tokens, result.flags) + if flag & tr.TokenFlag.SAMPLED + ] == [20, 21] + assert value.model_dump_json() == before diff --git a/tests/unit/trajectories/test_native_terminal.py b/tests/unit/trajectories/test_native_terminal.py index da2339027..c2b4eb885 100644 --- a/tests/unit/trajectories/test_native_terminal.py +++ b/tests/unit/trajectories/test_native_terminal.py @@ -251,3 +251,26 @@ def load(config: Any) -> Any: | tr.TokenFlag.STOP ) assert math.isnan(result.logprobs[boundary]) + + +def test_known_numeric_nonterminal_stop_keeps_terminal_length_offline(monkeypatch): + first = _chat_exchange([1], [2, 3]) + assert first.response.choices[0].model_extra is not None + first.response.choices[0].model_extra["stop_reason"] = 3 + last = _chat_exchange([1, 2, 3, 4], [5], offset=1) + last.response.choices[0].finish_reason = "length" + value = tr.Trajectory( + exchanges=tr.TrajectoryExchanges(chat_completions=[first, last]) + ) + + def no_load(*args, **kwargs): + raise AssertionError("Complete numeric-stop chain needs no tokenizer") + + monkeypatch.setattr(module, "_tokenizer_config", no_load) + monkeypatch.setattr(module, "_load_tokenizer", no_load) + result = value.tokenize() + assert result.tokens == [1, 2, 3, 4, 5] + assert result.flags[2] & tr.TokenFlag.STOP + assert not result.flags[-1] & tr.TokenFlag.STOP + assert result.logprobs[1:3] == [-0.2, -0.3] + assert result.logprobs[-1] == -0.5 diff --git a/tests/unit/trajectories/test_resolved_stop_authority.py b/tests/unit/trajectories/test_resolved_stop_authority.py index b32f6be19..e22a9c3f6 100644 --- a/tests/unit/trajectories/test_resolved_stop_authority.py +++ b/tests/unit/trajectories/test_resolved_stop_authority.py @@ -584,3 +584,50 @@ def test_context_snapshot_does_not_certify_cyclic_or_opaque_values(): module._tokenization_context(cyclic) with pytest.raises(TypeError, match="Unsupported mutable"): module._tokenization_context([object()]) + + +@pytest.mark.parametrize("kind", ["enum", "enum_state", "string"]) +@pytest.mark.parametrize("mutate", [False, True]) +def test_scalar_subclass_context_keeps_value_and_mutable_state(kind, mutate): + from enum import Enum + + class Content(str, Enum): + original = "public content" + changed = "changed content" + + class Text(str): + state: list[str] + + text = Text("public content") + text.state = ["original"] + selected = Content.original if kind in ("enum", "enum_state") else text + enum_state = ["original"] + if kind == "enum_state": + setattr(selected, "state", enum_state) + exchange = branch(0, length=False) + exchange.request["messages"][0]["content"] = selected + value = trajectory(exchange) + calls = [] + + class Tokenizer: + eos_token_id = 9 + + def convert_tokens_to_ids(self, token): + calls.append(token) + if mutate: + if kind == "enum": + exchange.request["messages"][0]["content"] = Content.changed + elif kind == "enum_state": + enum_state.append("changed") + else: + text.state.append("changed") + return None + + if mutate: + with pytest.raises(ValueError, match="context changed"): + value.tokenize(tokenizer=cast(Any, Tokenizer())) + else: + actual = value.tokenize(tokenizer=cast(Any, Tokenizer())) + assert actual.tokens + assert actual.flags[-1] & tr.TokenFlag.STOP + assert calls From 4ece40fe02ed6d68f1267a9d837dd365387fd1f4 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 20:25:35 +0000 Subject: [PATCH 20/67] Keep scoped reasoning normalization idempotent --- src/art_inference/chat_template.py | 10 ++++++++-- tests/unit/test_literal_reasoning_content.py | 13 ++++++++++++- 2 files changed, 20 insertions(+), 3 deletions(-) diff --git a/src/art_inference/chat_template.py b/src/art_inference/chat_template.py index adcfee50c..d35fbb33b 100644 --- a/src/art_inference/chat_template.py +++ b/src/art_inference/chat_template.py @@ -366,9 +366,15 @@ def chat_template_with_preserved_thinking(chat_template: object) -> object: "reasoning_content + '\\n\\n\\n'", "reasoning_content + ('\\n\\n' if preserve_thinking and message.reasoning_content is string and reasoning_content else '\\n\\n\\n')", ) - if not inline_parser_removed: + if ( + not inline_parser_removed + and "{%- set reasoning_content = reasoning_content|trim %}" + in literal_template + ): # The recognized parser path already preserved its own input. Do - # not apply the legacy template-wide trim rewrite to other macros. + # not apply the legacy trim rewrite without its recognized + # reasoning assignment: an unrelated inline filter can survive + # parser removal and must not activate this on a second call. chat_template = chat_template.replace( "set content = render_content(message.content, true)|trim", "set content = (render_content(message.content, true) if preserve_thinking and message.role == 'assistant' else render_content(message.content, true)|trim)", diff --git a/tests/unit/test_literal_reasoning_content.py b/tests/unit/test_literal_reasoning_content.py index 2a7f9bc30..26e5afb40 100644 --- a/tests/unit/test_literal_reasoning_content.py +++ b/tests/unit/test_literal_reasoning_content.py @@ -424,10 +424,15 @@ def test_plain_block_whitespace_keeps_prior_inline_parser_coverage(separator, qu @pytest.mark.parametrize("preserve", [False, True]) @pytest.mark.parametrize("newline", ["\n", "\r\n"]) @pytest.mark.parametrize("layout", ["macros", "same_line", "branches"]) -def test_content_trim_is_scoped_to_the_recognized_parser(preserve, newline, layout): +@pytest.mark.parametrize("inline_structured_reasoning", [False, True]) +def test_content_trim_is_scoped_to_the_recognized_parser( + preserve, newline, layout, inline_structured_reasoning +): match = _QWEN_INLINE_REASONING.search(_TEMPLATE) assert match is not None parser = match.group() + if inline_structured_reasoning: + parser += "{{ reasoning_content|trim }}" trim = "{% set content = render_content(message.content, true)|trim %}" render = "{% macro render_content(content, count) %}{{ content }}{% endmacro %}" preview = "{% macro preview(message) %}" + trim + "[{{ content }}]{% endmacro %}" @@ -465,6 +470,12 @@ def test_content_trim_is_scoped_to_the_recognized_parser(preserve, newline, layo ) ) assert chat_template_with_preserved_thinking(fixed) == fixed + assert ( + chat_template_with_preserved_thinking( + chat_template_with_preserved_thinking(fixed) + ) + == fixed + ) if layout == "branches": kwargs["message"] = {"role": "user", "content": " answer "} assert env.from_string(fixed).render(**kwargs).endswith("[answer]|[answer]") From 201fdf496d946f4ba2f78130a97ea69bb5ac3fa7 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 21:02:06 +0000 Subject: [PATCH 21/67] Retain consumed evidence and prove complete historical request context --- docs/features/additional-histories.mdx | 12 +- src/art/trajectories/_tokenize.py | 253 +++++++++++++++--- .../trajectories/test_consumed_evidence.py | 179 +++++++++++++ .../trajectories/test_literal_thinking_off.py | 6 +- .../test_named_template_arguments.py | 101 +++++++ .../test_responses_copied_stop_mutation.py | 102 +++++++ tests/unit/trajectories/test_tokenize.py | 7 +- 7 files changed, 612 insertions(+), 48 deletions(-) create mode 100644 tests/unit/trajectories/test_consumed_evidence.py create mode 100644 tests/unit/trajectories/test_named_template_arguments.py create mode 100644 tests/unit/trajectories/test_responses_copied_stop_mutation.py diff --git a/docs/features/additional-histories.mdx b/docs/features/additional-histories.mdx index 11bacd3d2..8befbf8ae 100644 --- a/docs/features/additional-histories.mdx +++ b/docs/features/additional-histories.mdx @@ -192,14 +192,16 @@ With a supplied renderer, request-owned assistant roles are proved throughout the native stream, including between sampled responses. This requirement applies even when template normalization makes no change. 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. When literal-content -normalization changes historical rendering, ART requires the original request -to reproduce the complete recorded prompt before recovering those roles. If -that proof fails, ART refuses to attach recorded logprobs to a changed prompt. +path, but cannot claim a complete rendered role mask. Recovering those roles +requires the original request to reproduce the complete recorded prompt, with +only independently proved whitespace retokenization permitted. Matching an +assistant body alone does not prove an altered role header. If that proof fails, +ART refuses to attach recorded logprobs to a changed prompt. 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 +STOP callback and checks all consumed evidence before returning, including +logprobs used only by the rendered fallback. Multi-history tokenization also checks completed histories after later renderer callbacks. Ordered original requests and history context, including roles, tools, render settings and protocol selectors, are captured before callbacks and checked diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index a11b79f33..2b064a57e 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -4,7 +4,7 @@ import codecs from collections.abc import Callable, Mapping, Sequence from copy import deepcopy -from dataclasses import dataclass, replace +from dataclasses import dataclass, field, replace from datetime import datetime from enum import Enum from functools import lru_cache @@ -596,6 +596,43 @@ def _translate_token_mask( return translated +def _recorded_prompt_tokens( + messages: list[dict[str, Any]], + *, + tokenizer: Tokenizer, + template: object, + tools: object, + kwargs: Mapping[str, object], +) -> list[int]: + messages, tools, kwargs = deepcopy((messages, tools, dict(kwargs))) + context = _render_context_key([messages, tools, kwargs]) + + def check() -> None: + if _render_context_key([messages, tools, kwargs]) != context: + raise ValueError( + "Renderer changed context while proving recorded request roles" + ) + + body = template + if isinstance(getattr(tokenizer, "chat_template", None), dict) and callable( + select := getattr(tokenizer, "get_chat_template", None) + ): + body = select( + chat_template=template if isinstance(template, str) else None, tools=tools + ) + check() + result = tokenizer.apply_chat_template( + normalize_tool_call_arguments_for_chat_template(messages, body), + tools=tools, + tokenize=True, + add_generation_prompt=True, + **({"chat_template": template} if template is not None else {}), + **kwargs, + ) + check() + return _ids(result) + + def _recorded_prompt_role_masks( messages: list[dict[str, Any]], sources: Sequence[object | None], @@ -959,6 +996,28 @@ class _TraceBuilder: validate_sources: Callable[[_SampledSourceKey | None], None] | None = None validate_context: Callable[[bool], None] | None = None track_sources: bool = True + consumed_sources: dict[_SampledSourceKey, object] = field(default_factory=dict) + + def consume_sources( + self, + sources: Mapping[_SampledSourceKey, object], + *, + selected_request_fields: tuple[str, ...] | None = None, + ) -> None: + # Keep every consumed source, including rendered-only logprobs that do + # not appear in the sampled-token trace. Validate before extending it. + if self.validate_sources is not None: + self.validate_sources(None) + added = { + key: source + for key, source in sources.items() + if key not in self.consumed_sources + } + if added or self.validate_sources is None: + self.consumed_sources.update(added) + self.validate_sources = _sampled_source_validator( + self.consumed_sources, selected_request_fields=selected_request_fields + ) def set( self, @@ -969,15 +1028,15 @@ def set( *, tokenizer: Tokenizer | None = None, ) -> None: - if self.validate_sources is not None: + if self.track_sources: + self.consume_sources(sources) + elif self.validate_sources is not None: self.validate_sources(None) self.tokenizer = tokenizer trace = _HistoryTokenizationTrace(source_keys=source_keys, sources=sources) trace.validate(tokenized) self.trace = trace self.rendered_outputs = rendered_outputs - if self.track_sources: - self.validate_sources = _sampled_source_validator(sources) def _fingerprint(value: object) -> str: @@ -2245,7 +2304,7 @@ def _response_message( def _resolved_chat_template( tokenizer: Tokenizer, template: object, tools: object -) -> tuple[object, dict[str, Any]]: +) -> tuple[object, 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) @@ -2261,7 +2320,7 @@ def _resolved_chat_template( if configured == selected: # apply_chat_template resolves names itself. Forwarding an # unchanged body could accidentally select a second named entry. - return template, defaults + return template, configured, defaults templates = getattr(tokenizer, "chat_template", None) if ( isinstance(configured, str) @@ -2272,7 +2331,7 @@ def _resolved_chat_template( "The normalized chat template is also a template name; " "cannot preserve the selected renderer without ambiguity" ) - return configured, defaults + return configured, configured, defaults def _template_ids( @@ -2322,13 +2381,17 @@ def _template_ids( or config.chat_template or getattr(tokenizer, "chat_template", None) ) - template, defaults = _resolved_chat_template(tokenizer, template, tools) + template, normalization_template, defaults = _resolved_chat_template( + tokenizer, template, tools + ) kwargs = { **defaults, **explicit_kwargs, } result = tokenizer.apply_chat_template( - normalize_tool_call_arguments_for_chat_template(messages, template), + normalize_tool_call_arguments_for_chat_template( + messages, normalization_template + ), tools=tools, tokenize=True, add_generation_prompt=not completed, @@ -2714,6 +2777,22 @@ def _tokenize_exchange_trajectory( selected_model = exchanges[0].model if selected_model is None: raise AssertionError("_exchange_list returned an exchange without a model") + consumed_sources = { + _exchange_sampled_source_key(exchange): exchange for exchange in exchanges + } + if _trace is not None: + _trace.consume_sources(consumed_sources) + assert _trace.validate_sources is not None + validate_consumed = _trace.validate_sources + else: + validate_consumed = _sampled_source_validator(consumed_sources) + + def checked(function: Callable[..., Any], *args: Any, **kwargs: Any) -> Any: + validate_consumed(None) + result = function(*args, **kwargs) + validate_consumed(None) + return result + exact_tokens = [_exchange_tokens(exchange) for exchange in exchanges] config = ( _TokenizerConfig(base_model if base_model is not None else selected_model) @@ -2740,7 +2819,7 @@ def _tokenize_exchange_trajectory( def fallback_config() -> _TokenizerConfig: nonlocal config if config is None: - config = _tokenizer_config(selected_model, base_model) + config = checked(_tokenizer_config, selected_model, base_model) return config for exchange, (prompt, completion, completion_logprobs) in zip( @@ -2800,8 +2879,9 @@ def fallback_config() -> _TokenizerConfig: if prompt is None: resolved_config = fallback_config() if tokenizer is None: - tokenizer = _load_tokenizer(resolved_config) - prompt = _template_ids( + tokenizer = checked(_load_tokenizer, resolved_config) + prompt = checked( + _template_ids, tokenizer, exchange, completed=False, @@ -2813,9 +2893,10 @@ def fallback_config() -> _TokenizerConfig: if completion is None: resolved_config = fallback_config() if tokenizer is None: - tokenizer = _load_tokenizer(resolved_config) + tokenizer = checked(_load_tokenizer, resolved_config) rendered_prompt = ( - _template_ids( + checked( + _template_ids, tokenizer, exchange, completed=False, @@ -2827,7 +2908,8 @@ def fallback_config() -> _TokenizerConfig: if prompt_is_exact else prompt ) - completed = _template_ids( + completed = checked( + _template_ids, tokenizer, exchange, completed=True, @@ -2841,8 +2923,8 @@ def fallback_config() -> _TokenizerConfig: "Completed response does not extend its generation prompt" ) completion = completed[len(rendered_prompt) :] - completion_logprobs = _align_visible_logprobs( - tokenizer, completion, exchange + completion_logprobs = checked( + _align_visible_logprobs, tokenizer, completion, exchange ) or [math.nan] * len(completion) if not token_ids: token_ids.extend(prompt) @@ -2854,8 +2936,9 @@ def fallback_config() -> _TokenizerConfig: elif len(prompt) < len(token_ids) or prompt[: len(token_ids)] != token_ids: resolved_config = fallback_config() if tokenizer is None: - tokenizer = _load_tokenizer(resolved_config) - repaired = _preserve_sampled_prefix( + tokenizer = checked(_load_tokenizer, resolved_config) + repaired = checked( + _preserve_sampled_prefix, prompt, token_ids, sampled_outputs, @@ -2866,7 +2949,8 @@ def fallback_config() -> _TokenizerConfig: raise ValueError( "Inference prompts do not form one append-only history" ) - current_render = _template_ids( + current_render = checked( + _template_ids, tokenizer, exchange, completed=False, @@ -2876,7 +2960,8 @@ def fallback_config() -> _TokenizerConfig: messages_override=messages_override, ) previous_exchange, previous_messages = previous_render_state - previous_render = _template_ids( + previous_render = checked( + _template_ids, tokenizer, previous_exchange, completed=True, @@ -2885,7 +2970,8 @@ def fallback_config() -> _TokenizerConfig: chat_template_kwargs=chat_template_kwargs, messages_override=previous_messages, ) - previous_canonical = _preserve_sampled_prefix( + previous_canonical = checked( + _preserve_sampled_prefix, previous_render, token_ids, sampled_outputs, @@ -2921,8 +3007,8 @@ def fallback_config() -> _TokenizerConfig: ) source_keys.extend([None] * len(suffix)) if len(completion_logprobs) != len(completion): - completion_logprobs = _align_visible_logprobs( - tokenizer, completion, exchange + completion_logprobs = checked( + _align_visible_logprobs, tokenizer, completion, exchange ) or [math.nan] * len(completion) token_ids.extend(completion) logprobs.extend(completion_logprobs) @@ -2957,6 +3043,7 @@ def fallback_config() -> _TokenizerConfig: sources, tokenizer=tokenizer, ) + validate_consumed(None) tokenized = TokenizedHistory( history=history, model=selected_model, @@ -4334,7 +4421,11 @@ def _source_has_no_materialized_output( return False -def _tokenization_context(value: object) -> object: +def _tokenization_context( + value: object, + *, + _observed: dict[int, tuple[object, object]] | None = None, +) -> object: """Snapshot ordered semantic inputs without serializing sampled responses. History and source fields remain typed, including protocol selectors and @@ -4343,7 +4434,7 @@ def _tokenization_context(value: object) -> object: """ # Shared message dictionaries occur in many recorded request prefixes. # Intern only within this one observation, never across callback checks. - observed: dict[int, tuple[object, object]] = {} + observed: dict[int, tuple[object, object]] = {} if _observed is None else _observed def snapshot(item: object) -> object: kind = type(item) @@ -4423,23 +4514,38 @@ def validate(require_supported: bool) -> None: def _sampled_source_validator( sources: Mapping[_SampledSourceKey, object], + *, + selected_request_fields: tuple[str, ...] | None = None, ) -> Callable[[_SampledSourceKey | None], None]: expected = {} + observed: dict[int, tuple[object, object]] = {} for key, source in sources.items(): exchange = _source_exchange(source) if exchange is None: raise ValueError("Sampled source has no exchange") + try: + request = _tokenization_context(exchange.request, _observed=observed) + except TypeError: + request = None expected[key] = ( source, exchange, exchange.model, _source_stop_evidence(source, key), - _tokenization_context_validator(exchange.request), + request, + _tokenization_context( + {name: exchange.request.get(name) for name in selected_request_fields} + ) + if selected_request_fields is not None + else None, ) def validate(selected: _SampledSourceKey | None) -> None: + # This observation is callback-free. Share aliased context containers + # across its sources, never across separate validations/callbacks. + observed: dict[int, tuple[object, object]] = {} for key in expected if selected is None else (selected,): - source, exchange, model, stop, validate_request = expected[key] + source, exchange, model, stop, request, selected_request = expected[key] current = ( _exchange_sampled_source_key(source) if isinstance(source, Exchange) @@ -4452,7 +4558,18 @@ def validate(selected: _SampledSourceKey | None) -> None: or _source_stop_evidence(source, key) != stop ): raise ValueError("Sampled source changed during tokenization callback") - validate_request(True) + if request is None: + raise ValueError("Tokenization context cannot be checked for callbacks") + current_request = exchange.request + if selected is not None and selected_request_fields is not None: + current_request = { + name: exchange.request.get(name) for name in selected_request_fields + } + request = selected_request + if _tokenization_context(current_request, _observed=observed) != request: + raise ValueError( + "Tokenization context changed during tokenization callback" + ) return validate @@ -5474,7 +5591,7 @@ def _tokenize_chat_view( if isinstance(tokenizer_template, str): template = tokenizer_template original_template = template - template, defaults = _resolved_chat_template( + template, normalization_template, defaults = _resolved_chat_template( resolved_tokenizer, template, history.tools ) kwargs = { @@ -5488,7 +5605,7 @@ def raw_render( selected_messages: list[dict[str, Any]], *, add_generation_prompt: bool ) -> list[int]: render_messages = normalize_tool_call_arguments_for_chat_template( - selected_messages, template + selected_messages, normalization_template ) return _ids( resolved_tokenizer.apply_chat_template( @@ -5521,7 +5638,7 @@ def render_text( ) -> str: return render_normalized_text( normalize_tool_call_arguments_for_chat_template( - selected_messages, template + selected_messages, normalization_template ), add_generation_prompt=add_generation_prompt, ) @@ -5549,11 +5666,18 @@ def render_text( for source in history.message_sources if id(source) in consumed_keys } - validate_consumed = _sampled_source_validator(consumed_sources) if _trace is not None: - if _trace.validate_sources is not None: - _trace.validate_sources(None) - _trace.validate_sources = validate_consumed + _trace.consume_sources( + consumed_sources, + selected_request_fields=("tools", "chat_template", "chat_template_kwargs"), + ) + assert _trace.validate_sources is not None + validate_consumed = _trace.validate_sources + else: + validate_consumed = _sampled_source_validator( + consumed_sources, + selected_request_fields=("tools", "chat_template", "chat_template_kwargs"), + ) def consumed_source_key(source: object) -> _SampledSourceKey: key = consumed_keys[id(source)] @@ -5580,7 +5704,7 @@ def segmented_render( # Normalization is message-local. Only the admitted nonmutating # renderer may share its normalized messages between prefixes. selected_messages = normalize_tool_call_arguments_for_chat_template( - selected_messages, template + selected_messages, normalization_template ) text = render_normalized_text( selected_messages, add_generation_prompt=add_generation_prompt @@ -5856,12 +5980,23 @@ 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 + needs_request_roles = any( + prior_message.get("role") == "assistant" + and ( + prior_source is None or not _source_is_sampled(prior_source) + ) + for prior_message, prior_source in zip( + history.messages[:message_index], + history.message_sources[:message_index], + strict=True, + ) + ) + if source_prompt != rendered_prompt and ( + original_template != template or needs_request_roles ): signature = _source_signature(source) request_context = None + full_prompt_proven = False try: # Canonical tool validation may reorder JSON keys. # Prove the historical prompt with the recorded @@ -5894,6 +6029,34 @@ def source_matches_context(source: object) -> bool: # Optional historical prefix rendering may be # unsupported although the selected renderer works. recorded_prompt_masks = None + if needs_request_roles and recorded_prompt_masks is None: + try: + raw_prompt = _recorded_prompt_tokens( + request_messages, + tokenizer=resolved_tokenizer, + template=original_template, + tools=request_tools, + kwargs=kwargs, + ) + # Whole-prompt correspondence is mandatory even + # if offsets/prefix probes are unsupported. The + # existing translator admits only exact IDs or + # proved whitespace retokenization, both ways. + _translate_token_mask( + raw_prompt, + source_prompt, + [True] * len(raw_prompt), + tokenizer=resolved_tokenizer, + ) + _translate_token_mask( + source_prompt, + raw_prompt, + [True] * len(source_prompt), + tokenizer=resolved_tokenizer, + ) + full_prompt_proven = True + except (TypeError, KeyError, NotImplementedError): + pass # These optional renderer calls must not turn cached # native evidence into authority for a changed source. prompt_cache.clear() @@ -5922,6 +6085,14 @@ def source_matches_context(source: object) -> bool: raise ValueError( "Sampled source changed while proving recorded request roles" ) + if ( + needs_request_roles + and recorded_prompt_masks is None + and not full_prompt_proven + ): + raise ValueError( + "Cannot preserve request roles without full recorded prompt proof" + ) break validate_consumed(None) @@ -6119,7 +6290,7 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: try: marked_text = resolved_tokenizer.apply_chat_template( normalize_tool_call_arguments_for_chat_template( - marked_messages, template + marked_messages, normalization_template ), tools=history.tools, tokenize=False, diff --git a/tests/unit/trajectories/test_consumed_evidence.py b/tests/unit/trajectories/test_consumed_evidence.py new file mode 100644 index 000000000..e7015b999 --- /dev/null +++ b/tests/unit/trajectories/test_consumed_evidence.py @@ -0,0 +1,179 @@ +from __future__ import annotations + +from typing import Any, cast + +from openai.types.chat import ChatCompletionMessageParam +import pytest +from test_tokenize import _CharacterTemplateTokenizer, _chat_exchange, _message_exchange + +import art.trajectories as tr + + +def trajectory(*exchanges): + return tr.Trajectory( + exchanges=tr.TrajectoryExchanges(chat_completions=list(exchanges)) + ) + + +@pytest.mark.parametrize("changed_header", [False, True]) +def test_request_role_requires_full_original_prompt(changed_header): + class Tokenizer(_CharacterTemplateTokenizer): + def apply_chat_template( + self, messages, *, tokenize=True, add_generation_prompt, **kwargs + ): + text = "".join(m["role"] + ":" + m["content"] + "§" for m in messages) + if add_generation_prompt: + text += "assistant:" + return self._encode(text) if tokenize else text + + tokenizer = Tokenizer() + messages = [ + {"role": "assistant", "content": "old answer"}, + {"role": "user", "content": "new question"}, + ] + text = tokenizer.apply_chat_template( + messages, tokenize=False, add_generation_prompt=True + ) + if changed_header: + text = text.replace("assistant:", "user_____:", 1) + prompt = tokenizer._encode(text) + exchange = _chat_exchange(prompt, tokenizer._encode("answer§")) + exchange.request["messages"] = cast(list[ChatCompletionMessageParam], messages) + value = trajectory(exchange) + if changed_header: + with pytest.raises(ValueError, match="[Rr]equest roles|assistant boundaries"): + value.tokenize(tokenizer=tokenizer) + else: + actual = value.tokenize(tokenizer=tokenizer) + assert actual.tokens[: len(prompt)] == prompt + assert actual.flags[len("assistant:")] & tr.TokenFlag.ASSISTANT + + +@pytest.mark.parametrize("mutate", [False, True]) +def test_exchange_fallback_binds_native_logprobs_before_prompt_render(mutate): + exchange = _message_exchange( + cast( + Any, + { + "model": "test/model", + "max_tokens": 5, + "messages": [{"role": "user", "content": "question"}], + }, + ), + token_ids=[2], + logprobs=[-0.2], + ) + calls = [] + + class Tokenizer: + def apply_chat_template(self, messages, **kwargs): + calls.append(True) + if mutate: + assert exchange.response.model_extra is not None + exchange.response.model_extra["logprobs"][0] = -9 + return [1] + + value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(messages=[exchange])) + if mutate: + with pytest.raises(ValueError, match="[Ss]ampled source changed"): + value.tokenize(tokenizer=cast(Any, Tokenizer())) + else: + actual = value.tokenize(tokenizer=cast(Any, Tokenizer())) + assert actual.tokens == [1, 2] and actual.logprobs[-1] == -0.2 + assert calls + + +@pytest.mark.parametrize("mutate", [False, True]) +def test_completed_rendered_logprobs_keep_consumed_validator(mutate): + first = _chat_exchange([1], [2]) + first.request["messages"] = [{"role": "user", "content": "first"}] + choice = first.response.choices[0] + assert ( + choice.model_extra is not None + and choice.logprobs is not None + and choice.logprobs.content is not None + ) + choice.model_extra.pop("token_ids") + choice.message.content = "a" + choice.logprobs.content[0].token = "a" + choice.logprobs.content[0].bytes = list(b"a") + second = _chat_exchange([3], [4], offset=1) + second.request["messages"] = [{"role": "user", "content": "second"}] + seen = [] + + class Tokenizer(_CharacterTemplateTokenizer): + def apply_chat_template(self, messages, **kwargs): + if messages and messages[0]["content"] == "second": + seen.append(True) + if mutate: + assert choice.logprobs is not None and choice.logprobs.content + choice.logprobs.content[0].logprob = -9 + return super().apply_chat_template(messages, **kwargs) + + value = trajectory(first, second) + if mutate: + with pytest.raises(ValueError, match="[Ss]ampled source changed"): + value.tokenize( + tokenizer=Tokenizer(), multi_history=True, chat_template="explicit" + ) + else: + actual = value.tokenize( + tokenizer=Tokenizer(), multi_history=True, chat_template="explicit" + ) + assert len(actual.histories) == 2 + assert -0.2 in actual.histories[0].logprobs + assert not any(f & tr.TokenFlag.SAMPLED for f in actual.histories[0].flags) + assert seen + + +@pytest.mark.parametrize("field", ["tools", "chat_template", "chat_template_kwargs"]) +@pytest.mark.parametrize("mutate", [False, True]) +def test_generic_evidence_reads_check_consumed_render_settings(field, mutate): + from test_evidence_reuse import extras + + second = None + calls = [] + + class Tokenizer(_CharacterTemplateTokenizer): + eos_token_id = None + all_special_tokens = [] + + def convert_tokens_to_ids(self, token): + return None + + def __call__(self, text, **kwargs): + if text in ("FIRST", "SECOND"): + calls.append(text) + assert second is not None + if mutate and text == "FIRST": + second.request[field] = { + "tools": [{"type": "function", "function": {"name": "edited"}}], + "chat_template": "edited", + "chat_template_kwargs": {"preserve_thinking": False}, + }[field] + elif mutate: + # A final-only check could miss this restoring callback. + second.request.pop(field, None) + return {"input_ids": [99]} + return super().__call__(text, **kwargs) + + tokenizer = Tokenizer() + prompt = tokenizer._encode("turn 0") + first = _chat_exchange(prompt, [20]) + second = _chat_exchange([*prompt, 20, *tokenizer._encode("turn 1")], [21], offset=1) + extras(first)["stop_reason"] = "FIRST" + extras(second)["stop_reason"] = "SECOND" + value = trajectory(first, second) + before = value.model_dump_json() + if mutate: + with pytest.raises(ValueError, match="[Cc]ontext changed"): + value.tokenize(tokenizer=tokenizer, chat_template="explicit override") + assert "FIRST" in calls and "SECOND" not in calls + else: + actual = value.tokenize(tokenizer=tokenizer, chat_template="explicit override") + assert [ + token + for token, flag in zip(actual.tokens, actual.flags, strict=True) + if flag & tr.TokenFlag.SAMPLED + ] == [20, 21] + assert value.model_dump_json() == before diff --git a/tests/unit/trajectories/test_literal_thinking_off.py b/tests/unit/trajectories/test_literal_thinking_off.py index f956c0b09..fda1d6aca 100644 --- a/tests/unit/trajectories/test_literal_thinking_off.py +++ b/tests/unit/trajectories/test_literal_thinking_off.py @@ -589,7 +589,11 @@ def test_named_template_selection_failure_keeps_original_error(): # 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, {}) + assert _tokenize._resolved_chat_template(tokenizer, custom, None) == ( + custom, + custom, + {}, + ) @pytest.mark.parametrize("selection", [None, "named"]) diff --git a/tests/unit/trajectories/test_named_template_arguments.py b/tests/unit/trajectories/test_named_template_arguments.py new file mode 100644 index 000000000..6c1e87c0d --- /dev/null +++ b/tests/unit/trajectories/test_named_template_arguments.py @@ -0,0 +1,101 @@ +from copy import deepcopy +from typing import Any, cast + +import pytest +from test_literal_thinking_off import _history, _NamedTemplateTokenizer + +import art.trajectories as tr +from art.trajectories import _tokenize +from art_inference.chat_template import configure_preserved_thinking_chat_template + + +@pytest.mark.parametrize("configured", [False, True]) +@pytest.mark.parametrize("selection", ["default", "named", "tools"]) +@pytest.mark.parametrize("route", ["history", "prompt", "completed"]) +def test_selected_body_normalizes_tool_arguments_without_reselecting( + configured, selection, route +): + tokenizer = _NamedTemplateTokenizer() + tokenizer.chat_template = { + name: body.replace("tool_call.arguments|items", "tool_call.arguments.items()") + for name, body in tokenizer.chat_template.items() + } + if configured: + configure_preserved_thinking_chat_template(tokenizer) + # A body can itself be another name: forwarding it would select WRONG. + selected = tokenizer.chat_template[ + "tool_use" if selection == "tools" else selection + ] + if configured: + tokenizer.chat_template[selected] = "WRONG" + templates = deepcopy(tokenizer.chat_template) + selector = "named" if selection == "named" else None + tools: list[Any] | None = ( + [{"type": "function", "function": {"name": "lookup", "parameters": {}}}] + if selection == "tools" + else None + ) + messages = [ + {"role": "user", "content": "Public query."}, + { + "role": "assistant", + "content": "", + "tool_calls": [ + { + "id": "public-call", + "type": "function", + "function": { + "name": "lookup", + "arguments": '{"city":"Paris","count":2}', + }, + } + ], + }, + ] + kwargs = {"enable_thinking": True, "preserve_thinking": False} + if route == "history": + history = tr.ChatCompletionsHistory( + model="public/qwen35", + messages=messages, + message_sources=[None, None], + tools=tools, + ) + before = history.model_dump() + result = history.tokenize( + tokenizer=tokenizer, + chat_template=selector, + chat_template_kwargs=kwargs, + ) + rendered = tokenizer.decode(result.tokens) + assert not any(flag & tr.TokenFlag.SAMPLED for flag in result.flags) + assert history.model_dump() == before + else: + original, _ = _history() + source = original.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["messages"] = cast(Any, messages) + if tools is not None: + exchange.request["tools"] = tools + before = exchange.model_dump() + result = _tokenize._template_ids( + tokenizer, + exchange, + completed=route == "completed", + config=_tokenize._TokenizerConfig(base_model="public/qwen35"), + chat_template=selector, + chat_template_kwargs=kwargs, + ) + rendered = tokenizer.decode(result) + assert exchange.model_dump() == before + assert "\nParis\n" in rendered + assert "\n2\n" in rendered + assert "WRONG" not in rendered + assert tokenizer.chat_template == templates + assert tokenizer.selected + assert all( + settings["enable_thinking"] is True and settings["preserve_thinking"] is False + for settings in tokenizer.settings + ) diff --git a/tests/unit/trajectories/test_responses_copied_stop_mutation.py b/tests/unit/trajectories/test_responses_copied_stop_mutation.py new file mode 100644 index 000000000..d4a961744 --- /dev/null +++ b/tests/unit/trajectories/test_responses_copied_stop_mutation.py @@ -0,0 +1,102 @@ +from __future__ import annotations + +import copy +import math +from typing import Any + +from openai.types.responses import Response +import pytest +from test_tokenize import _response_exchange + +import art.trajectories as tr +from art.trajectories import _tokenize as module + + +@pytest.mark.parametrize("mutate", [False, True]) +def test_copied_responses_stop_cannot_return_stale_prior_logprobs( + monkeypatch: pytest.MonkeyPatch, mutate: bool +) -> None: + exchange = _response_exchange("three-generations", 2, prompt_token_ids=[1]) + payload = exchange.response.model_dump(mode="python") + message = payload["output"][0] + payload["output"] = [] + for index in range(3): + item = copy.deepcopy(message) + item["id"] = f"public-message-{index}" + item["content"][0]["text"] = f"answer{index}" + payload["output"].append(item) + payload["token_generations"] = [ + { + "prompt_token_ids": prompt, + "output_tokens": [ + {"token_id": token, "logprob": -token / 10} for token in output + ], + "output_indices": [index], + } + for index, (prompt, output) in enumerate( + [([1], [2]), ([1, 2, 3], [4, 5]), ([1, 2, 3, 5, 6], [7])] + ) + ] + exchange.response = Response.model_validate(payload) + trajectory = tr.Trajectory(exchanges=tr.TrajectoryExchanges(responses=[exchange])) + before = trajectory.model_dump_json() + + class Tokenizer: + armed = False + copied_callback = False + + def convert_tokens_to_ids(self, token: str) -> None: + del token + if self.armed: + self.armed = False + self.copied_callback = True + if mutate: + assert exchange.response.model_extra is not None + exchange.response.model_extra["token_generations"][0][ + "output_tokens" + ][0]["logprob"] = -9.5 + + def apply_chat_template(self, *args: Any, **kwargs: Any) -> Any: + raise AssertionError("complete native records must not render") + + tokenizer: Any = Tokenizer() + native = module._tokenize_exact_responses_history + + def observe(history: Any, **kwargs: Any) -> Any: + # Arm only after preflight and entry to the history containing generation + # 2. Its first converter callback is the copied generation-1 STOP probe. + # This spy preserves the real callback and avoids global call ordinals. + tokenizer.armed = any( + source is not None and source.generation_index == 2 + for source in history.input_sources + ) + return native(history, **kwargs) + + monkeypatch.setattr(module, "_tokenize_exact_responses_history", observe) + if mutate: + with pytest.raises(ValueError, match="Sampled source changed"): + trajectory.tokenize(tokenizer=tokenizer, multi_history=True) + assert tokenizer.copied_callback + assert exchange.response.model_extra is not None + assert ( + exchange.response.model_extra["token_generations"][0]["output_tokens"][0][ + "logprob" + ] + == -9.5 + ) + else: + result = trajectory.tokenize(tokenizer=tokenizer, multi_history=True) + assert tokenizer.copied_callback + assert [history.tokens for history in result.histories] == [ + [1, 2, 3, 4, 5], + [1, 2, 3, 5, 6, 7], + ] + original, copied = result.histories + assert original.logprobs[1] == copied.logprobs[1] == -0.2 + assert original.logprobs[3:] == [-0.4, -0.5] + assert copied.logprobs[-1] == -0.7 + assert math.isnan(copied.logprobs[3]) + assert copied.flags[3] == ( + tr.TokenFlag.EXACT | tr.TokenFlag.ASSISTANT | tr.TokenFlag.OUTPUT + ) + assert trajectory.model_dump_json() == before diff --git a/tests/unit/trajectories/test_tokenize.py b/tests/unit/trajectories/test_tokenize.py index fc86d06c3..6af6d9688 100644 --- a/tests/unit/trajectories/test_tokenize.py +++ b/tests/unit/trajectories/test_tokenize.py @@ -7539,7 +7539,12 @@ def apply_chat_template( assert len(datums) == 4 -def test_responses_prompt_repair_opt_in_uses_native_text_and_source_position() -> None: +def test_responses_prompt_repair_opt_in_uses_native_text_and_source_position( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr( + "art.trajectories._tokenize._WARNED_PREFIX_RETOKENIZATION", False + ) exchange = _response_exchange("repeated-retokenization", 101) data = exchange.response.model_dump(mode="python") data["output"].append( From 1466a4e17d1fdd7e4a2ab3c021254ea03722bde1 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 21:05:28 +0000 Subject: [PATCH 22/67] Check consumed evidence around prefix repair warning hooks --- src/art/trajectories/_tokenize.py | 2 +- .../trajectories/test_consumed_evidence.py | 70 +++++++++++++++++++ 2 files changed, 71 insertions(+), 1 deletion(-) diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index 2b064a57e..f00d580e9 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -2990,7 +2990,7 @@ def fallback_config() -> _TokenizerConfig: ] prompt = repaired prompt_is_exact = False - _warn_prefix_retokenization() + checked(_warn_prefix_retokenization) suffix = prompt[len(token_ids) :] token_ids.extend(suffix) logprobs.extend([math.nan] * len(suffix)) diff --git a/tests/unit/trajectories/test_consumed_evidence.py b/tests/unit/trajectories/test_consumed_evidence.py index e7015b999..fb7589cdf 100644 --- a/tests/unit/trajectories/test_consumed_evidence.py +++ b/tests/unit/trajectories/test_consumed_evidence.py @@ -177,3 +177,73 @@ def __call__(self, text, **kwargs): if flag & tr.TokenFlag.SAMPLED ] == [20, 21] assert value.model_dump_json() == before + + +@pytest.mark.parametrize("mutate", [False, True]) +def test_exchange_warning_hook_cannot_change_consumed_logprobs(monkeypatch, mutate): + import warnings + + monkeypatch.setattr( + "art.trajectories._tokenize._WARNED_PREFIX_RETOKENIZATION", False + ) + first = _message_exchange( + cast( + Any, + { + "model": "test/model", + "max_tokens": 5, + "messages": [{"role": "user", "content": "q"}], + }, + ), + prompt_token_ids=[1], + token_ids=[2], + logprobs=[-0.2], + ) + second = _message_exchange( + cast( + Any, + { + "model": "test/model", + "max_tokens": 5, + "messages": [ + {"role": "user", "content": "q"}, + { + "role": "assistant", + "content": [{"type": "text", "text": "answer"}], + }, + {"role": "user", "content": "next"}, + ], + }, + ), + identifier="second", + offset=1, + prompt_token_ids=[1, 3, 4], + token_ids=[5], + logprobs=[-0.5], + ) + value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(messages=[first, second])) + history = value.anthropic_messages_history( + reconcile_text_equivalent_tokenizations=True + ) + calls = [] + + def showwarning(*args, **kwargs): + calls.append(True) + if mutate: + assert second.response.model_extra is not None + second.response.model_extra["logprobs"][0] = -9 + + class Tokenizer: + def __call__(self, text, **kwargs): + assert text == "answer" + return [3] + + monkeypatch.setattr(warnings, "showwarning", showwarning) + if mutate: + with pytest.raises(ValueError, match="[Ss]ampled source changed"): + history.tokenize(tokenizer=cast(Any, Tokenizer())) + else: + actual = history.tokenize(tokenizer=cast(Any, Tokenizer())) + assert actual.tokens == [1, 2, 4, 5] + assert actual.logprobs[1] == -0.2 and actual.logprobs[-1] == -0.5 + assert calls From f5bd90b1cfcbf6bf8788cc2a7e4ae5324dcb8aab Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 21:17:26 +0000 Subject: [PATCH 23/67] Retain distinct consumed sources and allow unused opaque native context --- src/art/trajectories/_tokenize.py | 177 ++++++++++++------ .../test_consumed_source_identity.py | 170 +++++++++++++++++ 2 files changed, 288 insertions(+), 59 deletions(-) create mode 100644 tests/unit/trajectories/test_consumed_source_identity.py diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index f00d580e9..37d7670bb 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 Callable, Mapping, Sequence +from collections.abc import Callable, Iterable, Mapping, Sequence from copy import deepcopy from dataclasses import dataclass, field, replace from datetime import datetime @@ -988,35 +988,56 @@ def _require_causal_predecessor(trainable: Sequence[bool]) -> None: raise ValueError("A trainable trajectory cannot start with a sampled token") +class _SampledSourceValidator(Protocol): + def __call__( + self, + selected: _SampledSourceKey | None, + *, + require_supported_context: bool = True, + ) -> None: ... + + @dataclass 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 + validate_sources: _SampledSourceValidator | None = None validate_context: Callable[[bool], None] | None = None track_sources: bool = True - consumed_sources: dict[_SampledSourceKey, object] = field(default_factory=dict) + consumed_sources: dict[ + tuple[_SampledSourceKey, int], tuple[_SampledSourceKey, object] + ] = field(default_factory=dict) def consume_sources( self, - sources: Mapping[_SampledSourceKey, object], + sources: Mapping[_SampledSourceKey, object] + | Sequence[tuple[_SampledSourceKey, object]], *, selected_request_fields: tuple[str, ...] | None = None, + require_supported_context: bool = True, ) -> None: - # Keep every consumed source, including rendered-only logprobs that do - # not appear in the sampled-token trace. Validate before extending it. + # Semantic keys survive protocol copies. Retain every consumed object, + # including aliases and rendered-only sources absent from the trace. if self.validate_sources is not None: - self.validate_sources(None) + self.validate_sources( + None, require_supported_context=require_supported_context + ) + items: Iterable[tuple[_SampledSourceKey, object]] = ( + cast(Mapping[_SampledSourceKey, object], sources).items() + if isinstance(sources, Mapping) + else sources + ) added = { - key: source - for key, source in sources.items() - if key not in self.consumed_sources + (key, id(source)): (key, source) + for key, source in items + if (key, id(source)) not in self.consumed_sources } if added or self.validate_sources is None: self.consumed_sources.update(added) self.validate_sources = _sampled_source_validator( - self.consumed_sources, selected_request_fields=selected_request_fields + list(self.consumed_sources.values()), + selected_request_fields=selected_request_fields, ) def set( @@ -1029,9 +1050,11 @@ def set( tokenizer: Tokenizer | None = None, ) -> None: if self.track_sources: - self.consume_sources(sources) + self.consume_sources( + sources, require_supported_context=tokenizer is not None + ) elif self.validate_sources is not None: - self.validate_sources(None) + self.validate_sources(None, require_supported_context=tokenizer is not None) self.tokenizer = tokenizer trace = _HistoryTokenizationTrace(source_keys=source_keys, sources=sources) trace.validate(tokenized) @@ -2777,17 +2800,24 @@ def _tokenize_exchange_trajectory( selected_model = exchanges[0].model if selected_model is None: raise AssertionError("_exchange_list returned an exchange without a model") - consumed_sources = { - _exchange_sampled_source_key(exchange): exchange for exchange in exchanges - } + consumed_sources = [ + (_exchange_sampled_source_key(exchange), exchange) for exchange in exchanges + ] if _trace is not None: - _trace.consume_sources(consumed_sources) + _trace.consume_sources( + consumed_sources, + require_supported_context=tokenizer_instance is not None, + ) assert _trace.validate_sources is not None validate_consumed = _trace.validate_sources else: validate_consumed = _sampled_source_validator(consumed_sources) + callback_used = False + def checked(function: Callable[..., Any], *args: Any, **kwargs: Any) -> Any: + nonlocal callback_used + callback_used = True validate_consumed(None) result = function(*args, **kwargs) validate_consumed(None) @@ -3043,7 +3073,9 @@ def fallback_config() -> _TokenizerConfig: sources, tokenizer=tokenizer, ) - validate_consumed(None) + validate_consumed( + None, require_supported_context=callback_used or tokenizer is not None + ) tokenized = TokenizedHistory( history=history, model=selected_model, @@ -4513,13 +4545,19 @@ def validate(require_supported: bool) -> None: def _sampled_source_validator( - sources: Mapping[_SampledSourceKey, object], + sources: Mapping[_SampledSourceKey, object] + | Sequence[tuple[_SampledSourceKey, object]], *, selected_request_fields: tuple[str, ...] | None = None, -) -> Callable[[_SampledSourceKey | None], None]: +) -> _SampledSourceValidator: expected = {} observed: dict[int, tuple[object, object]] = {} - for key, source in sources.items(): + items: Iterable[tuple[_SampledSourceKey, object]] = ( + cast(Mapping[_SampledSourceKey, object], sources).items() + if isinstance(sources, Mapping) + else sources + ) + for key, source in items: exchange = _source_exchange(source) if exchange is None: raise ValueError("Sampled source has no exchange") @@ -4527,49 +4565,70 @@ def _sampled_source_validator( request = _tokenization_context(exchange.request, _observed=observed) except TypeError: request = None - expected[key] = ( - source, - exchange, - exchange.model, - _source_stop_evidence(source, key), - request, - _tokenization_context( - {name: exchange.request.get(name) for name in selected_request_fields} + expected.setdefault(key, []).append( + ( + source, + exchange, + exchange.model, + _source_stop_evidence(source, key), + request, + _tokenization_context( + { + name: exchange.request.get(name) + for name in selected_request_fields + } + ) + if selected_request_fields is not None + else None, ) - if selected_request_fields is not None - else None, ) - def validate(selected: _SampledSourceKey | None) -> None: + def validate( + selected: _SampledSourceKey | None, + *, + require_supported_context: bool = True, + ) -> None: # This observation is callback-free. Share aliased context containers # across its sources, never across separate validations/callbacks. observed: dict[int, tuple[object, object]] = {} for key in expected if selected is None else (selected,): - source, exchange, model, stop, request, selected_request = 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") - if request is None: - raise ValueError("Tokenization context cannot be checked for callbacks") - current_request = exchange.request - if selected is not None and selected_request_fields is not None: - current_request = { - name: exchange.request.get(name) for name in selected_request_fields - } - request = selected_request - if _tokenization_context(current_request, _observed=observed) != request: - raise ValueError( - "Tokenization context changed during tokenization callback" + for source, exchange, model, stop, request, selected_request in 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" + ) + if request is None: + if require_supported_context: + raise ValueError( + "Tokenization context cannot be checked for callbacks" + ) + continue + current_request = exchange.request + if selected is not None and selected_request_fields is not None: + current_request = { + name: exchange.request.get(name) + for name in selected_request_fields + } + request = selected_request + if ( + _tokenization_context(current_request, _observed=observed) + != request + ): + raise ValueError( + "Tokenization context changed during tokenization callback" + ) return validate @@ -5661,11 +5720,11 @@ def render_text( for source in history.message_sources if source is not None and _source_is_sampled(source) } - consumed_sources = { - consumed_keys[id(source)]: source + consumed_sources = [ + (consumed_keys[id(source)], source) for source in history.message_sources if id(source) in consumed_keys - } + ] if _trace is not None: _trace.consume_sources( consumed_sources, diff --git a/tests/unit/trajectories/test_consumed_source_identity.py b/tests/unit/trajectories/test_consumed_source_identity.py new file mode 100644 index 000000000..80f46c5b6 --- /dev/null +++ b/tests/unit/trajectories/test_consumed_source_identity.py @@ -0,0 +1,170 @@ +from __future__ import annotations + +import copy +from typing import Any, cast + +import pytest +from test_tokenize import _chat_exchange, _completion_exchange, _message_exchange + +import art.trajectories as tr +import art.trajectories._tokenize as core + + +def test_complete_native_messages_accept_opaque_unused_metadata(): + exchange = _message_exchange( + cast( + Any, + { + "model": "test/model", + "max_tokens": 5, + "messages": [{"role": "user", "content": "question"}], + }, + ), + prompt_token_ids=[1], + token_ids=[2], + logprobs=[-0.2], + ) + opaque = object() + exchange.request["metadata"] = cast(Any, {"opaque": opaque}) + value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(messages=[exchange])) + actual = value.tokenize() + assert actual.tokens == [1, 2] + assert actual.logprobs[-1] == -0.2 + assert actual.flags[-1] & tr.TokenFlag.SAMPLED + assert exchange.request["metadata"]["opaque"] is opaque + + +@pytest.mark.parametrize("mutate", [False, True]) +def test_consumed_equal_key_distinct_source_objects_remain_guarded(mutate): + first = _chat_exchange([1], [2]) + second = copy.deepcopy(first) + key = core._exchange_sampled_source_key(first) + assert key == core._exchange_sampled_source_key(second) + assert second is not first + builder = core._TraceBuilder() + builder.consume_sources({key: first}) + builder.consume_sources({key: second}) + if mutate: + assert second.response.choices[0].logprobs is not None + assert second.response.choices[0].logprobs.content + second.response.choices[0].logprobs.content[0].logprob = -9.0 + assert builder.validate_sources is not None + if mutate: + with pytest.raises(ValueError, match="[Ss]ampled source changed"): + builder.validate_sources(None) + else: + builder.validate_sources(None) + + +def test_complete_native_completions_accept_opaque_unused_metadata(): + exchange = _completion_exchange() + opaque = object() + exchange.request["metadata"] = cast(Any, {"opaque": opaque}) + value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(completions=[exchange])) + actual = value.tokenize() + assert actual.tokens == [1, 2] + assert actual.logprobs[-1] == -0.2 + assert actual.flags[-1] & tr.TokenFlag.SAMPLED + assert exchange.request["metadata"]["opaque"] is opaque + + +def test_opaque_exchange_still_refuses_before_prompt_callback(): + exchange = _message_exchange( + cast( + Any, + { + "model": "test/model", + "max_tokens": 5, + "messages": [{"role": "user", "content": "question"}], + }, + ), + token_ids=[2], + logprobs=[-0.2], + ) + exchange.request["metadata"] = cast(Any, {"opaque": object()}) + calls = [] + + class Tokenizer: + def apply_chat_template(self, messages, **kwargs): + calls.append(True) + return [1] + + value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(messages=[exchange])) + with pytest.raises(ValueError, match="context cannot be checked"): + value.tokenize(tokenizer=cast(Any, Tokenizer())) + assert not calls + + +@pytest.mark.parametrize("selected", [False, True]) +@pytest.mark.parametrize("which", [0, 1]) +@pytest.mark.parametrize("field", ["logprob", "request"]) +def test_aliases_in_one_consumption_keep_each_original(selected, which, field): + first = _chat_exchange([1], [2]) + second = copy.deepcopy(first) + key = core._exchange_sampled_source_key(first) + builder = core._TraceBuilder() + builder.consume_sources([(key, first), (key, second)]) + source = (first, second)[which] + if field == "request": + source.request["messages"][0]["content"] = "edited" + else: + assert source.response.choices[0].logprobs is not None + assert source.response.choices[0].logprobs.content + source.response.choices[0].logprobs.content[0].logprob = -9.0 + assert builder.validate_sources is not None + with pytest.raises( + ValueError, match="[Ss]ampled source changed|[Cc]ontext changed" + ): + builder.validate_sources(key if selected else None) + + +def test_alias_extension_does_not_bless_earlier_mutation(): + first = _chat_exchange([1], [2]) + second = copy.deepcopy(first) + key = core._exchange_sampled_source_key(first) + builder = core._TraceBuilder() + builder.consume_sources({key: first}) + assert first.response.choices[0].logprobs is not None + assert first.response.choices[0].logprobs.content + first.response.choices[0].logprobs.content[0].logprob = -9.0 + with pytest.raises(ValueError, match="[Ss]ampled source changed"): + builder.consume_sources({key: second}) + + +def test_opaque_allowance_never_disables_native_evidence_validation(): + exchange = _chat_exchange([1], [2]) + exchange.request["metadata"] = cast(Any, {"opaque": object()}) + key = core._exchange_sampled_source_key(exchange) + validate = core._sampled_source_validator({key: exchange}) + validate(None, require_supported_context=False) + assert exchange.response.choices[0].logprobs is not None + assert exchange.response.choices[0].logprobs.content + exchange.response.choices[0].logprobs.content[0].logprob = -9.0 + with pytest.raises(ValueError, match="[Ss]ampled source changed"): + validate(None, require_supported_context=False) + + +def test_callback_free_multiple_histories_keep_opaque_metadata(): + first = _completion_exchange() + second = copy.deepcopy(first) + second.request["model"] = "other/model" + opaque = object() + for exchange in (first, second): + exchange.request["metadata"] = cast(Any, {"opaque": opaque}) + value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(completions=[first, second])) + actual = value.tokenize(multi_history=True) + assert len(actual.histories) == 2 + assert [history.tokens for history in actual.histories] == [[1, 2], [1, 2]] + assert [history.logprobs[-1] for history in actual.histories] == [-0.2, -0.2] + assert all(history.flags[-1] & tr.TokenFlag.SAMPLED for history in actual.histories) + assert first.request["metadata"]["opaque"] is opaque + assert second.request["metadata"]["opaque"] is opaque + + +def test_supported_context_is_checked_without_opaque_requirement(): + source = _chat_exchange([1], [2]) + key = core._exchange_sampled_source_key(source) + validate = core._sampled_source_validator({key: source}) + source.request["messages"][0]["content"] = "edited" + with pytest.raises(ValueError, match="[Cc]ontext changed"): + validate(None, require_supported_context=False) From 1033b83a94e5a0f9fd07f0004e37e84451d23922 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 21:21:10 +0000 Subject: [PATCH 24/67] Preserve content bindings across conditional parser joins --- src/art_inference/chat_template.py | 76 +++++++++------- .../test_reasoning_parser_joined_branches.py | 89 +++++++++++++++++++ 2 files changed, 132 insertions(+), 33 deletions(-) create mode 100644 tests/unit/test_reasoning_parser_joined_branches.py diff --git a/src/art_inference/chat_template.py b/src/art_inference/chat_template.py index d35fbb33b..3c85c0e06 100644 --- a/src/art_inference/chat_template.py +++ b/src/art_inference/chat_template.py @@ -65,9 +65,14 @@ def visit_Output(self, node: nodes.Output, *args: Any, **kwargs: Any): def operations(text: str): return WithoutWhitespace().visit(env.parse(text)).body - operation = operations( - "".join("{% " + statement + " %}" for statement in _QWEN_INLINE_STATEMENTS) - ) + parser = "".join("{% " + statement + " %}" for statement in _QWEN_INLINE_STATEMENTS) + operation = operations(parser) + structured_operation = operations( + "{% if message.reasoning_content is string %}" + "{% set reasoning_content = message.reasoning_content %}{% else %}" + + parser + + "{% endif %}" + )[0] normalized = re.sub(r"\r\n?", "\n", template) offsets = [ i @@ -210,50 +215,55 @@ def assistant_condition(test: nodes.Node) -> bool | None: return equal if test.ops[0].op == "eq" else not equal return None - def visit(body: Sequence[nodes.Node], binding: nodes.Assign | None = None) -> None: + def visit(body: Sequence[nodes.Node], bindings: set[int]) -> set[int]: + bindings = bindings.copy() for node in body: if isinstance(node, nodes.If): - if id(node) in edited_parsers and node == operation[0]: - if binding is not None: - selected.add(id(binding)) - binding = None - continue - if binding is not None and reads_content(node.test): - shared.add(id(binding)) - condition = assistant_condition(node.test) - if condition is not False: - visit(node.body, binding) - if condition is not True: - # elif_ is a list of If nodes whose else is stored on the - # outer If. Stop following it once a role match is proved. - for branch in node.elif_: - visit([branch], binding) - if assistant_condition(branch.test) is True: - break - else: - visit(node.else_, binding) - if any( - writes_content(n) - for n in node.find_all((nodes.Assign, nodes.AssignBlock)) + if (id(node) in edited_parsers and node == operation[0]) or ( + node == structured_operation + and any(id(child) in edited_parsers for child in node.else_) ): - binding = None + # The recognized structured/inline reasoning selector is + # one normalization boundary too, preserving its existing + # literal-content contract in either reasoning mode. + if len(bindings) == 1: + selected.update(bindings) + else: + shared.update(bindings) # No unique consumed assignment. + bindings.clear() + continue + joined = set() + # If does not introduce a Jinja scope. Retain every binding + # reaching the join, including paths that skipped the parser. + for branch in (node, *node.elif_): + if reads_content(branch.test): + shared.update(bindings) + condition = assistant_condition(branch.test) + if condition is not False: + joined.update(visit(branch.body, bindings)) + if condition is True: + break + else: + joined.update(visit(node.else_, bindings)) + bindings = joined else: - if binding is not None and reads_content(node): - shared.add(id(binding)) + if reads_content(node): + shared.update(bindings) if isinstance(node, nodes.Assign): if writes_content(node): - binding = node + bindings = {id(node)} else: # Macro/loop/with/block bodies have independent bindings. for _, value in node.iter_fields(): if isinstance(value, list) and all( isinstance(n, nodes.Node) for n in value ): - visit(value) + visit(value, set()) if isinstance(node, nodes.AssignBlock) and writes_content(node): - binding = None + bindings.clear() + return bindings - visit(tree.body) + visit(tree.body, set()) trims = [ env.parse("{% set content = " + content + " %}").body[0] for content in ( diff --git a/tests/unit/test_reasoning_parser_joined_branches.py b/tests/unit/test_reasoning_parser_joined_branches.py new file mode 100644 index 000000000..06a6b0f8b --- /dev/null +++ b/tests/unit/test_reasoning_parser_joined_branches.py @@ -0,0 +1,89 @@ +from jinja2.sandbox import ImmutableSandboxedEnvironment +import pytest +from test_literal_reasoning_content import _TEMPLATE + +from art_inference.chat_template import ( + _QWEN_INLINE_REASONING, + _without_inline_reasoning_parser, +) + + +@pytest.mark.parametrize("scope", ["top", "macro"]) +@pytest.mark.parametrize("layout", ["if", "elif", "nested", "sequential"]) +@pytest.mark.parametrize("content", [" answer ", " beforexafter "]) +def test_joined_preview_retains_its_original_trim(scope, layout, content): + match = _QWEN_INLINE_REASONING.search(_TEMPLATE) + assert match is not None + parser = match.group() + trim = "{% set content = render_content(message.content, true)|trim %}" + if layout == "if": + branch = "{% if not preview_only %}" + parser + "{% endif %}" + elif layout == "elif": + branch = ( + "{% if preview_only %}{% elif mode == 'answer' %}" + parser + "{% endif %}" + ) + elif layout == "nested": + branch = ( + "{% if outer %}{% if not preview_only %}" + + parser + + "{% endif %}{% endif %}" + ) + else: + branch = ( + "{% if not preview_only %}" + parser + "{% endif %}" + "{% if mode == 'other' %}" + parser + "{% endif %}" + ) + body = trim + branch + "[{{ content }}]" + if scope == "macro": + body = ( + "{% macro answer(message) %}" + body + "{% endmacro %}{{ answer(message) }}" + ) + template = ( + "{% macro render_content(content, count) %}{{ content }}{% endmacro %}" + body + ) + fixed = _without_inline_reasoning_parser(template) + env = ImmutableSandboxedEnvironment(trim_blocks=True, lstrip_blocks=True) + kwargs = dict( + message={"role": "assistant", "content": content}, + preview_only=True, + mode="answer", + outer=True, + ) + assert env.from_string(fixed).render(**kwargs) == env.from_string(template).render( + **kwargs + ) + assert env.from_string(fixed).render(**kwargs) == "[" + content.strip() + "]" + assert trim in fixed + assert not _QWEN_INLINE_REASONING.search(fixed) + # The parser path still treats the text as literal; its shared trim stays. + kwargs["preview_only"] = False + assert env.from_string(fixed).render(**kwargs) == "[" + content.strip() + "]" + assert _without_inline_reasoning_parser(fixed) == fixed + + +@pytest.mark.parametrize("mode", ["a", "b", "c"]) +def test_all_joined_paths_consuming_parser_can_preserve_assistant_whitespace(mode): + match = _QWEN_INLINE_REASONING.search(_TEMPLATE) + assert match is not None + parser = match.group() + template = ( + "{% macro render_content(content, count) %}{{ content }}{% endmacro %}" + "{% set content = render_content(message.content, true)|trim %}" + "{% if mode == 'a' %}" + + parser + + "{% elif mode == 'b' %}" + + parser + + "{% else %}" + + parser + + "{% endif %}[{{ content }}]" + ) + fixed = _without_inline_reasoning_parser(template) + content = " beforexafter " + env = ImmutableSandboxedEnvironment(trim_blocks=True, lstrip_blocks=True) + assert ( + env.from_string(fixed).render( + mode=mode, message={"role": "assistant", "content": content} + ) + == "[" + content + "]" + ) + assert _without_inline_reasoning_parser(fixed) == fixed From c27a555b6b28e34faf0029ef5ffb52d560c41a7c Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 21:39:02 +0000 Subject: [PATCH 25/67] Bind consumed protocol evidence before tokenizer callbacks --- docs/features/additional-histories.mdx | 5 + src/art/trajectories/_tokenize.py | 293 ++++++++++++++---- .../test_protocol_consumed_evidence.py | 210 +++++++++++++ .../test_responses_copied_stop_mutation.py | 18 +- .../test_visible_response_logprobs.py | 94 ++++++ 5 files changed, 552 insertions(+), 68 deletions(-) create mode 100644 tests/unit/trajectories/test_protocol_consumed_evidence.py create mode 100644 tests/unit/trajectories/test_visible_response_logprobs.py diff --git a/docs/features/additional-histories.mdx b/docs/features/additional-histories.mdx index 8befbf8ae..86c0f583d 100644 --- a/docs/features/additional-histories.mdx +++ b/docs/features/additional-histories.mdx @@ -203,12 +203,17 @@ source tokens, logprobs, model or stop evidence. ART checks each source around i STOP callback and checks all consumed evidence before returning, including logprobs used only by the rendered fallback. Multi-history tokenization also checks completed histories after later renderer callbacks. +Evidence is bound when first read, including a later Responses generation read +to prove a copied suffix, visible item logprobs, and text used for prefix repair. +It is not a snapshot of every future response before that response is used. Ordered original requests and history context, including roles, tools, render settings and protocol selectors, are captured before callbacks and checked before reuse. Unexpected callback exceptions keep their original identity. These checks refuse stale results; they do not make callbacks or source objects immutable. Context snapshots reuse shared containers only within one observation; callback-free native assembly retains its bounded response-evidence reuse. +Unused opaque request metadata remains valid on complete callback-free native +paths; callbacks require context that ART can compare. 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 37d7670bb..7649e0c92 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -1005,6 +1005,10 @@ class _TraceBuilder: validate_sources: _SampledSourceValidator | None = None validate_context: Callable[[bool], None] | None = None track_sources: bool = True + rendered_evidence: bool = False + auxiliary_evidence: dict[object, tuple[Callable[[], object], object]] = field( + default_factory=dict + ) consumed_sources: dict[ tuple[_SampledSourceKey, int], tuple[_SampledSourceKey, object] ] = field(default_factory=dict) @@ -1016,6 +1020,7 @@ def consume_sources( *, selected_request_fields: tuple[str, ...] | None = None, require_supported_context: bool = True, + rendered_evidence: bool = False, ) -> None: # Semantic keys survive protocol copies. Retain every consumed object, # including aliases and rendered-only sources absent from the trace. @@ -1033,13 +1038,49 @@ def consume_sources( for key, source in items if (key, id(source)) not in self.consumed_sources } - if added or self.validate_sources is None: + if ( + added + or self.validate_sources is None + or rendered_evidence + and not self.rendered_evidence + ): self.consumed_sources.update(added) + self.rendered_evidence |= rendered_evidence self.validate_sources = _sampled_source_validator( list(self.consumed_sources.values()), selected_request_fields=selected_request_fields, + rendered_evidence=self.rendered_evidence, ) + def consume_auxiliary( + self, identity: object, read: Callable[[], object], value: object + ) -> None: + if identity in self.auxiliary_evidence: + self.validate_auxiliary() + else: + self.auxiliary_evidence[identity] = read, _tokenization_context(value) + + def validate_auxiliary(self) -> None: + for read, expected in self.auxiliary_evidence.values(): + if _tokenization_context(read()) != expected: + raise ValueError( + "Consumed source text changed during tokenization callback" + ) + + def checked(self, function: Callable[..., Any], *args: Any, **kwargs: Any) -> Any: + if self.validate_context is not None: + self.validate_context(True) + if self.validate_sources is not None: + self.validate_sources(None) + self.validate_auxiliary() + result = function(*args, **kwargs) + if self.validate_sources is not None: + self.validate_sources(None) + if self.validate_context is not None: + self.validate_context(True) + self.validate_auxiliary() + return result + def set( self, tokenized: TokenizedHistory, @@ -2803,25 +2844,18 @@ def _tokenize_exchange_trajectory( consumed_sources = [ (_exchange_sampled_source_key(exchange), exchange) for exchange in exchanges ] - if _trace is not None: - _trace.consume_sources( - consumed_sources, - require_supported_context=tokenizer_instance is not None, - ) - assert _trace.validate_sources is not None - validate_consumed = _trace.validate_sources - else: - validate_consumed = _sampled_source_validator(consumed_sources) - + ledger = _trace or _TraceBuilder(track_sources=False) + ledger.consume_sources( + consumed_sources, require_supported_context=tokenizer_instance is not None + ) callback_used = False def checked(function: Callable[..., Any], *args: Any, **kwargs: Any) -> Any: nonlocal callback_used + if not callback_used: + ledger.consume_sources([], rendered_evidence=True) callback_used = True - validate_consumed(None) - result = function(*args, **kwargs) - validate_consumed(None) - return result + return ledger.checked(function, *args, **kwargs) exact_tokens = [_exchange_tokens(exchange) for exchange in exchanges] config = ( @@ -3073,7 +3107,8 @@ def fallback_config() -> _TokenizerConfig: sources, tokenizer=tokenizer, ) - validate_consumed( + assert ledger.validate_sources is not None + ledger.validate_sources( None, require_supported_context=callback_used or tokenizer is not None ) tokenized = TokenizedHistory( @@ -3980,7 +4015,54 @@ def _tokenize_exact_responses_history( source_keys: list[_SampledSourceKey | None] = [] sources: dict[_SampledSourceKey, object] = {} sampled_outputs: list[_SampledOutput] = [] + ledger = _trace or _TraceBuilder(track_sources=False) + pending: list[tuple[_SampledSourceKey, object]] = [] + + def checked( + function: Callable[..., Any], *args: Any, _callbacks: bool = True, **kwargs: Any + ) -> Any: + # Keys bind records at first use. Before a callout, register the records + # read since the last callback-free boundary; never reuse past a callout. + if _callbacks: + ledger.consume_sources(pending) + pending.clear() + return ( + ledger.checked(function, *args, **kwargs) + if _callbacks + else function(*args, **kwargs) + ) + + observed: dict[tuple[int, int], tuple[_SampledSourceKey, object]] = {} + + def decline() -> None: + # A later rendered fallback must retain the evidence already inspected + # here, even if this pure native attempt could not assemble a stream. + if pending: + ledger.consume_sources(pending, require_supported_context=False) + pending.clear() + + def observe( + exchange: ResponsesExchange, generation_index: int + ) -> tuple[_SampledSourceKey, object]: + identity = id(exchange), generation_index + if identity not in observed: + 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 + ) + observed[identity] = _sampled_source_key(source), source + pending.append(observed[identity]) + key, source = observed[identity] + if (key, id(source)) in ledger.consumed_sources: + assert ledger.validate_sources is not None + ledger.validate_sources(key) + return key, source + for position, (exchange, generation_index) in enumerate(generation_keys): + source_key, source = observe(exchange, generation_index) generations = _response_generations(exchange.response) if not 0 <= generation_index < len(generations): raise ValueError("Responses source generation index is out of bounds") @@ -3988,19 +4070,13 @@ def _tokenize_exact_responses_history( prompt = generation.prompt_token_ids 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 - ) + return decline() context_only = False retained = retained_output_indices.get((id(exchange), generation_index), set()) following_prompt = None if position + 1 < len(generation_keys): following_exchange, following_index = generation_keys[position + 1] + observe(following_exchange, following_index) following_prompt = _response_generations(following_exchange.response)[ following_index ].prompt_token_ids @@ -4013,14 +4089,14 @@ def _tokenize_exact_responses_history( ) if retained != set(generation.output_indices) or copied_suffix: if position + 1 >= len(generation_keys): - return None + return decline() next_exchange, next_generation_index = generation_keys[position + 1] next_generations = _response_generations(next_exchange.response) if not 0 <= next_generation_index < len(next_generations): raise ValueError("Responses source generation index is out of bounds") next_prompt = next_generations[next_generation_index].prompt_token_ids if next_prompt is None: - return None + return decline() retained_suffix = _retained_output_suffix( prompt=prompt, output=output, @@ -4028,7 +4104,7 @@ def _tokenize_exact_responses_history( later_prompt=next_prompt, ) if retained_suffix is None: - return None + return decline() context_only = retained_suffix[0] != output if context_only and not _complete_source_is_represented( source, prompt, output, generation.output_logprobs, _prior @@ -4042,9 +4118,17 @@ def _tokenize_exact_responses_history( output_text = None else: output_logprobs = generation.output_logprobs - output_text = generation.output_text or _response_generation_text( - exchange.response, generation - ) + + def read_text( + exchange: ResponsesExchange = exchange, + generation: _ResponseGeneration = generation, + ) -> str | None: + return generation.output_text or _response_generation_text( + exchange.response, generation + ) + + output_text = read_text() + ledger.consume_auxiliary((source_key, id(source)), read_text, output_text) if not token_ids: token_ids.extend(prompt) logprobs.extend([math.nan] * len(prompt)) @@ -4059,11 +4143,18 @@ def _tokenize_exact_responses_history( source_keys.extend([None] * len(suffix)) else: if tokenizer is None and sampled_outputs: - tokenizer = _load_tokenizer( - _tokenizer_config(history.model, base_model) + tokenizer = checked( + _load_tokenizer, + checked(_tokenizer_config, history.model, base_model), ) repaired = ( - _preserve_sampled_prefix(prompt, token_ids, sampled_outputs, tokenizer) + checked( + _preserve_sampled_prefix, + prompt, + token_ids, + sampled_outputs, + tokenizer, + ) if tokenizer is not None else None ) @@ -4071,7 +4162,7 @@ def _tokenize_exact_responses_history( raise ValueError( "Responses token generations do not form one append-only history" ) - _warn_prefix_retokenization() + checked(_warn_prefix_retokenization) suffix = repaired[len(token_ids) :] token_ids.extend(suffix) logprobs.extend([math.nan] * len(suffix)) @@ -4088,27 +4179,16 @@ def _tokenize_exact_responses_history( ] * len(output) ) - 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 - ), - None, - ) - if source is None: - raise AssertionError("Responses generation has no history source") - source_key = _sampled_source_key(source) source_keys.extend([None if context_only else source_key] * len(output)) sources[source_key] = source if context_only: - stop_count = _sampled_stop_suffix( + stop_count = checked( + _sampled_stop_suffix, generation.output_token_ids or [], source=source, source_key=source_key, tokenizer=tokenizer, + _callbacks=tokenizer is not None, ) for offset in range( max(len(token_ids) - len(output), len(token_ids) - stop_count), @@ -4123,13 +4203,19 @@ def _tokenize_exact_responses_history( start=len(token_ids) - len(output), ) ) - _mark_sampled_stops( + checked( + _mark_sampled_stops, token_ids, flags, source_keys, sources, tokenizer=tokenizer, + _callbacks=tokenizer is not None, ) + if pending and (ledger.track_sources or ledger.validate_sources is not None): + ledger.consume_sources(pending, require_supported_context=tokenizer is not None) + if ledger.validate_sources is not None: + ledger.validate_sources(None, require_supported_context=tokenizer is not None) tokenized = TokenizedHistory( history=history, model=history.model, @@ -4549,6 +4635,7 @@ def _sampled_source_validator( | Sequence[tuple[_SampledSourceKey, object]], *, selected_request_fields: tuple[str, ...] | None = None, + rendered_evidence: bool = False, ) -> _SampledSourceValidator: expected = {} observed: dict[int, tuple[object, object]] = {} @@ -4580,6 +4667,14 @@ def _sampled_source_validator( ) if selected_request_fields is not None else None, + _tokenization_context( + _visible_logprobs( + exchange, + source=None if isinstance(source, Exchange) else source, + ) + ) + if rendered_evidence and isinstance(exchange, ResponsesExchange) + else None, ) ) @@ -4592,9 +4687,15 @@ def validate( # across its sources, never across separate validations/callbacks. observed: dict[int, tuple[object, object]] = {} for key in expected if selected is None else (selected,): - for source, exchange, model, stop, request, selected_request in expected[ - key - ]: + for ( + source, + exchange, + model, + stop, + request, + selected_request, + visible, + ) in expected[key]: current = ( _exchange_sampled_source_key(source) if isinstance(source, Exchange) @@ -4609,6 +4710,19 @@ def validate( raise ValueError( "Sampled source changed during tokenization callback" ) + if ( + visible is not None + and _tokenization_context( + _visible_logprobs( + exchange, + source=None if isinstance(source, Exchange) else source, + ) + ) + != visible + ): + raise ValueError( + "Rendered source logprobs changed during tokenization callback" + ) if request is None: if require_supported_context: raise ValueError( @@ -5729,6 +5843,7 @@ def render_text( _trace.consume_sources( consumed_sources, selected_request_fields=("tools", "chat_template", "chat_template_kwargs"), + rendered_evidence=True, ) assert _trace.validate_sources is not None validate_consumed = _trace.validate_sources @@ -5736,6 +5851,7 @@ def render_text( validate_consumed = _sampled_source_validator( consumed_sources, selected_request_fields=("tools", "chat_template", "chat_template_kwargs"), + rendered_evidence=True, ) def consumed_source_key(source: object) -> _SampledSourceKey: @@ -7622,9 +7738,14 @@ def _tokenize_completions_token_history( flags[start:end] = [TokenFlag.EXACT | TokenFlag.SAMPLED | TokenFlag.OUTPUT] * ( end - start ) + ledger = _trace or _TraceBuilder(track_sources=False) + consumed: list[tuple[_SampledSourceKey, object]] = [] for span in history.prompt_sources: if span.source is None: continue + if tokenizer is not None or ledger.track_sources: + evidence_source = _completion_evidence_source(span.source) + consumed.append((_sampled_source_key(evidence_source), evidence_source)) if span.source.choice_index is not None: source_key = _sampled_source_key(span.source) source_keys[span.start : span.end] = [source_key] * (span.end - span.start) @@ -7650,6 +7771,10 @@ def _tokenize_completions_token_history( continue if len(selected_logprobs) == span.end - span.start: logprobs[span.start : span.end] = selected_logprobs + if consumed: + ledger.consume_sources( + consumed, require_supported_context=tokenizer is not None + ) _mark_sampled_stops( history.prompt, flags, @@ -7657,6 +7782,8 @@ def _tokenize_completions_token_history( sources, tokenizer=tokenizer, ) + if ledger.validate_sources is not None: + ledger.validate_sources(None, require_supported_context=tokenizer is not None) tokenized = TokenizedHistory( history=history, model=history.model, @@ -7746,13 +7873,30 @@ def _tokenize_completions_string_history( raise ValueError("Completions sampled spans are out of bounds") sampled[start:end] = [True] * (end - start) + ledger = _trace or _TraceBuilder(track_sources=False) + pending: list[tuple[_SampledSourceKey, object]] = [] + + def checked( + function: Callable[..., Any], *args: Any, _callbacks: bool = True, **kwargs: Any + ) -> Any: + # Keys bind records at first use. Before a callout, register the records + # read since the last callback-free boundary; never reuse past a callout. + if _callbacks: + ledger.consume_sources(pending) + pending.clear() + return ( + ledger.checked(function, *args, **kwargs) + if _callbacks + else function(*args, **kwargs) + ) + config: _TokenizerConfig | None = None def resolved_tokenizer() -> Tokenizer: nonlocal config, tokenizer if tokenizer is None: - config = config or _tokenizer_config(history.model, base_model) - tokenizer = _load_tokenizer(config) + config = config or checked(_tokenizer_config, history.model, base_model) + tokenizer = checked(_load_tokenizer, config) return tokenizer token_ids: list[int] = [] @@ -7767,6 +7911,8 @@ def resolved_tokenizer() -> Tokenizer: source_logprobs: list[float] = [] is_sampled = any(sampled[span.start : span.end]) if source is not None: + evidence_source = _completion_evidence_source(source) + pending.append((_sampled_source_key(evidence_source), evidence_source)) prompt, completion, prompt_logprobs, completion_logprobs = ( _completion_source_evidence(source) ) @@ -7814,14 +7960,20 @@ def resolved_tokenizer() -> Tokenizer: ids = ( exact if exact is not None - else _ids(resolved_tokenizer()(text, add_special_tokens=False)) + else _ids(checked(resolved_tokenizer(), text, add_special_tokens=False)) ) token_ids.extend(ids) if exact is not None and len(source_logprobs) == len(ids): logprobs.extend(source_logprobs) else: visible = ( - _completion_visible_logprobs(source, text, resolved_tokenizer(), ids) + checked( + _completion_visible_logprobs, + source, + text, + resolved_tokenizer(), + ids, + ) if source is not None and source.choice_index is not None else None ) @@ -7840,13 +7992,19 @@ def resolved_tokenizer() -> Tokenizer: sources[source_key] = source else: source_keys.extend([None] * len(ids)) - _mark_sampled_stops( + checked( + _mark_sampled_stops, token_ids, flags, source_keys, sources, tokenizer=tokenizer, + _callbacks=tokenizer is not None, ) + if pending and (ledger.track_sources or ledger.validate_sources is not None): + ledger.consume_sources(pending, require_supported_context=tokenizer is not None) + if ledger.validate_sources is not None: + ledger.validate_sources(None, require_supported_context=tokenizer is not None) tokenized = TokenizedHistory( history=history, model=history.model, @@ -7859,9 +8017,7 @@ def resolved_tokenizer() -> Tokenizer: return tokenized -def _completion_source_evidence( - source: CompletionsSource, -) -> tuple[list[int] | None, list[int] | None, list[float], list[float]]: +def _completion_evidence_source(source: CompletionsSource) -> CompletionsSource: from ._history import _completion_choice_groups prompt_groups = _completion_choice_groups(source.exchange) @@ -7889,6 +8045,22 @@ def _completion_source_evidence( ) if selected is None: raise ValueError("Completions choice source does not belong to its prompt") + return ( + source + if source.choice_index is not None + else source.model_copy(update={"choice_index": selected.index}) + ) + + +def _completion_source_evidence( + source: CompletionsSource, +) -> tuple[list[int] | None, list[int] | None, list[float], list[float]]: + selected_source = _completion_evidence_source(source) + selected = next( + choice + for choice in source.exchange.response.choices + if choice.index == selected_source.choice_index + ) return _completion_evidence( source.exchange.response.model_copy(update={"choices": [selected]}), echo=source.exchange.request.get("echo") is True, @@ -8269,6 +8441,7 @@ def _validate_completed_sources(builders: Sequence[_TraceBuilder | None]) -> Non builder.validate_context(True) if builder.validate_sources is not None: builder.validate_sources(None) + builder.validate_auxiliary() def _complete_resolved_sampled_stops( diff --git a/tests/unit/trajectories/test_protocol_consumed_evidence.py b/tests/unit/trajectories/test_protocol_consumed_evidence.py new file mode 100644 index 000000000..305f7c12b --- /dev/null +++ b/tests/unit/trajectories/test_protocol_consumed_evidence.py @@ -0,0 +1,210 @@ +from __future__ import annotations + +from typing import Any, cast + +from openai.types.responses import Response, ResponseOutputMessage, ResponseOutputText +import pytest +from test_tokenize import ( + _completion_exchange, + _multi_output_responses_chat_history, + _response_exchange, +) + +import art.trajectories as tr +import art.trajectories._tokenize as core + + +@pytest.mark.parametrize("mutate", [False, True]) +@pytest.mark.parametrize("target", ["logprob", "text"]) +def test_responses_prefix_repair_retains_consumed_current_logprobs( + monkeypatch, mutate, target +): + monkeypatch.setattr(core, "_WARNED_PREFIX_RETOKENIZATION", False) + exchange = _response_exchange("repair-consumption", 2, prompt_token_ids=[1]) + data = exchange.response.model_dump(mode="python") + data["output"].append( + { + "id": "second", + "type": "message", + "role": "assistant", + "status": "completed", + "content": [ + { + "type": "output_text", + "text": "last", + "annotations": [], + "logprobs": [], + } + ], + } + ) + data["token_generations"] = [ + { + "prompt_token_ids": [1], + "output_tokens": [{"token_id": 2, "logprob": -0.2}], + "output_indices": [0], + }, + { + "prompt_token_ids": [1, 3, 4], + "output_tokens": [{"token_id": 5, "logprob": -0.5}], + "output_indices": [1], + }, + ] + exchange.response = Response.model_validate(data) + value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(responses=[exchange])) + history = value.responses_history(reconcile_text_equivalent_tokenizations=True) + before = value.model_dump_json() + calls = [] + + class Tokenizer: + def __call__(self, text, **kwargs): + calls.append(text) + if mutate and text == "answer": + if target == "text": + output = exchange.response.output[0] + assert isinstance(output, ResponseOutputMessage) + content = output.content[0] + assert isinstance(content, ResponseOutputText) + content.text = "edited" + else: + assert exchange.response.model_extra is not None + exchange.response.model_extra["token_generations"][1][ + "output_tokens" + ][0]["logprob"] = -9.0 + return {"answer": [3], "last": [5]}[text] + + def apply_chat_template(self, *args, **kwargs): + raise AssertionError("native repair must not render") + + if mutate: + with pytest.raises( + ValueError, match="[Ss]ampled source changed|Consumed source text changed" + ): + history.tokenize(tokenizer=cast(Any, Tokenizer())) + else: + with pytest.warns(UserWarning, match="preserved the original sampled"): + result = history.tokenize(tokenizer=cast(Any, Tokenizer())) + assert result.tokens == [1, 2, 4, 5] + assert result.logprobs[1] == -0.2 and result.logprobs[-1] == -0.5 + assert result.flags[-1] & tr.TokenFlag.SAMPLED + assert value.model_dump_json() == before + assert "answer" in calls + + +@pytest.mark.parametrize("mutate", [False, True]) +def test_completions_rendered_only_logprobs_bind_before_encoding(monkeypatch, mutate): + exchange = _completion_exchange(prompt="q") + choice = exchange.response.choices[0] + assert choice.model_extra is not None and choice.logprobs is not None + choice.model_extra.pop("prompt_token_ids") + choice.model_extra.pop("token_ids") + choice.text = "a" + choice.logprobs.tokens = ["a"] + choice.logprobs.token_logprobs = [-0.2] + value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(completions=[exchange])) + before = value.model_dump_json() + calls = [] + armed = False + original = core._completion_visible_logprobs + + def observe(*args, **kwargs): + nonlocal armed + armed = True + try: + return original(*args, **kwargs) + finally: + armed = False + + monkeypatch.setattr(core, "_completion_visible_logprobs", observe) + + class Tokenizer: + def __call__(self, text, **kwargs): + if armed: + calls.append(text) + if mutate: + assert ( + choice.logprobs is not None and choice.logprobs.token_logprobs + ) + choice.logprobs.token_logprobs[0] = -9.0 + return {"q": [1], "a": [2]}[text] + + if mutate: + with pytest.raises(ValueError, match="[Ss]ampled source changed"): + value.tokenize(tokenizer=cast(Any, Tokenizer())) + else: + result = value.tokenize(tokenizer=cast(Any, Tokenizer())) + assert result.tokens == [1, 2] and result.logprobs[-1] == -0.2 + assert result.flags[-1] & tr.TokenFlag.OUTPUT + assert not result.flags[-1] & tr.TokenFlag.SAMPLED + assert value.model_dump_json() == before + assert calls == ["a"] + + +def test_new_ledger_callback_preserves_raised_object_after_mutation(): + exchange = _completion_exchange() + key = core._exchange_sampled_source_key(exchange) + builder = core._TraceBuilder() + builder.consume_sources({key: exchange}) + error = RuntimeError("public callback failure") + + def callback(): + choice = exchange.response.choices[0] + assert choice.logprobs is not None and choice.logprobs.token_logprobs + choice.logprobs.token_logprobs[0] = -9.0 + raise error + + with pytest.raises(RuntimeError) as caught: + builder.checked(callback) + assert caught.value is error + + +@pytest.mark.parametrize("protocol", ["responses", "completions"]) +def test_native_standalone_protocol_needs_no_callback_validator(monkeypatch, protocol): + def unexpected(*args, **kwargs): + raise AssertionError( + "callback-free standalone native must not build callback guards" + ) + + monkeypatch.setattr(core, "_sampled_source_validator", unexpected) + if protocol == "responses": + exchange = _response_exchange("offline", 2, prompt_token_ids=[1]) + value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(responses=[exchange])) + expected_lp = -0.1 + else: + exchange = _completion_exchange() + value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(completions=[exchange])) + expected_lp = -0.2 + actual = value.tokenize() + assert actual.tokens == [1, 2] and actual.logprobs[-1] == expected_lp + assert actual.flags[-1] & tr.TokenFlag.SAMPLED + + +def test_declined_native_responses_keeps_original_consumed_key(): + projected = _multi_output_responses_chat_history() + source = next( + source + for source in projected.message_sources + if source is not None and source.generation_index == 0 + ) + exchange = source.exchange + assert isinstance(exchange, tr.ResponsesExchange) + value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(responses=[exchange])) + history = value.responses_history() + # One retained item cannot stand for the complete two-item native generation. + history.input = history.input[:-1] + history.input_sources = history.input_sources[:-1] + builder = core._TraceBuilder(track_sources=False) + assert ( + core._tokenize_exact_responses_history( + history, base_model=None, tokenizer=None, _trace=builder + ) + is None + ) + assert builder.validate_sources is not None + builder.validate_sources(None, require_supported_context=False) + assert exchange.response.model_extra is not None + exchange.response.model_extra["token_generations"][0]["output_tokens"][0][ + "logprob" + ] = -9.0 + with pytest.raises(ValueError, match="[Ss]ampled source changed"): + builder.validate_sources(None, require_supported_context=False) diff --git a/tests/unit/trajectories/test_responses_copied_stop_mutation.py b/tests/unit/trajectories/test_responses_copied_stop_mutation.py index d4a961744..5ad54ce54 100644 --- a/tests/unit/trajectories/test_responses_copied_stop_mutation.py +++ b/tests/unit/trajectories/test_responses_copied_stop_mutation.py @@ -13,8 +13,9 @@ @pytest.mark.parametrize("mutate", [False, True]) +@pytest.mark.parametrize("changed_generation", [0, 2]) def test_copied_responses_stop_cannot_return_stale_prior_logprobs( - monkeypatch: pytest.MonkeyPatch, mutate: bool + monkeypatch: pytest.MonkeyPatch, mutate: bool, changed_generation: int ) -> None: exchange = _response_exchange("three-generations", 2, prompt_token_ids=[1]) payload = exchange.response.model_dump(mode="python") @@ -52,9 +53,9 @@ def convert_tokens_to_ids(self, token: str) -> None: self.copied_callback = True if mutate: assert exchange.response.model_extra is not None - exchange.response.model_extra["token_generations"][0][ - "output_tokens" - ][0]["logprob"] = -9.5 + exchange.response.model_extra["token_generations"][ + changed_generation + ]["output_tokens"][0]["logprob"] = -9.5 def apply_chat_template(self, *args: Any, **kwargs: Any) -> Any: raise AssertionError("complete native records must not render") @@ -65,7 +66,8 @@ def apply_chat_template(self, *args: Any, **kwargs: Any) -> Any: def observe(history: Any, **kwargs: Any) -> Any: # Arm only after preflight and entry to the history containing generation # 2. Its first converter callback is the copied generation-1 STOP probe. - # This spy preserves the real callback and avoids global call ordinals. + # Generation 2's prompt has already been read as a suffix witness, so + # that record is consumed too. Preserve real callbacks without ordinals. tokenizer.armed = any( source is not None and source.generation_index == 2 for source in history.input_sources @@ -79,9 +81,9 @@ def observe(history: Any, **kwargs: Any) -> Any: assert tokenizer.copied_callback assert exchange.response.model_extra is not None assert ( - exchange.response.model_extra["token_generations"][0]["output_tokens"][0][ - "logprob" - ] + exchange.response.model_extra["token_generations"][changed_generation][ + "output_tokens" + ][0]["logprob"] == -9.5 ) else: diff --git a/tests/unit/trajectories/test_visible_response_logprobs.py b/tests/unit/trajectories/test_visible_response_logprobs.py new file mode 100644 index 000000000..e073256e6 --- /dev/null +++ b/tests/unit/trajectories/test_visible_response_logprobs.py @@ -0,0 +1,94 @@ +"""Public, explicit-rendering Responses evidence callback regression.""" + +from __future__ import annotations + +import math +from typing import Any + +from openai.types.responses import Response, ResponseOutputMessage, ResponseOutputText +import pytest +from test_tokenize import _multi_output_responses_chat_history + +import art.trajectories as tr + + +@pytest.mark.parametrize("route", ["history", "trajectory"]) +@pytest.mark.parametrize("mutate", [False, True]) +def test_multi_output_generation_keeps_consumed_visible_logprobs(route, mutate): + projected = _multi_output_responses_chat_history() + source = next( + source + for source in projected.message_sources + if source is not None and source.output_indices == (0,) + ) + exchange = source.exchange + assert isinstance(exchange, tr.ResponsesExchange) + data = exchange.response.model_dump(mode="python") + # Retain one exact generation spanning both outputs: neither individual + # rendered message owns all of its native output evidence. + assert data["token_generations"][0]["output_indices"] == [0, 1] + for output, text, lp in zip( + data["output"], ["first", "second"], [-0.1, -0.2], strict=True + ): + output["content"][0]["logprobs"] = [ + { + "token": text, + "logprob": lp, + "bytes": list(text.encode()), + "top_logprobs": [], + } + ] + exchange.response = Response.model_validate(data) + value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(responses=[exchange])) + output = exchange.response.output[0] + assert isinstance(output, ResponseOutputMessage) + content = output.content[0] + assert isinstance(content, ResponseOutputText) and content.logprobs + first_lp = content.logprobs[0] + calls = [] + + class Tokenizer: + def __call__(self, text: str, **kwargs: Any): + token = {"turn 0": 1, "first": 20, "second": 30}[text] + if kwargs.get("return_offsets_mapping"): + calls.append(text) + if mutate and text == "first": + first_lp.logprob = -9.0 + return {"input_ids": [token], "offset_mapping": [(0, len(text))]} + return [token] + + def apply_chat_template(self, messages, **kwargs): + return [token for message in messages for token in self(message["content"])] + + def tokenize(): + if route == "history": + return value.responses_history().tokenize( + tokenizer=Tokenizer(), chat_template="custom" + ) + return value.tokenize(tokenizer=Tokenizer(), chat_template="custom") + + if mutate: + with pytest.raises( + ValueError, + match="[Ss]ampled source changed|[Cc]onsumed|[Cc]ontext changed|Rendered source logprobs changed", + ): + actual = tokenize() + # If the guard misses the edit, prove this is the consumed old + # value being returned, not a benign pre-consumption refresh. + assert first_lp.logprob == -9.0 + assert actual.tokens == [1, 20, 30] + assert actual.logprobs[1:] == [-0.1, -0.2] + assert not any(flag & tr.TokenFlag.SAMPLED for flag in actual.flags) + assert first_lp.logprob == -9.0 + else: + actual = tokenize() + assert actual.tokens == [1, 20, 30] + assert math.isnan(actual.logprobs[0]) + assert actual.logprobs[1:] == [-0.1, -0.2] + assert actual.flags == [ + tr.TokenFlag(0), + tr.TokenFlag.ASSISTANT | tr.TokenFlag.OUTPUT, + tr.TokenFlag.ASSISTANT | tr.TokenFlag.OUTPUT, + ] + assert first_lp.logprob == -0.1 + assert "first" in calls From e11885a68385851ca2ef45b2fb149ebad9cc4997 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 22:45:26 +0000 Subject: [PATCH 26/67] Guard rendered token evidence and callback inputs before reuse --- docs/features/additional-histories.mdx | 11 +- src/art/trajectories/_tokenize.py | 433 +++++++++++++++--- .../test_interior_request_prompt.py | 78 ++++ .../test_renderer_context_callbacks.py | 410 +++++++++++++++++ .../test_response_item_token_id.py | 112 +++++ .../test_visible_response_logprobs.py | 2 +- 6 files changed, 970 insertions(+), 76 deletions(-) create mode 100644 tests/unit/trajectories/test_interior_request_prompt.py create mode 100644 tests/unit/trajectories/test_renderer_context_callbacks.py create mode 100644 tests/unit/trajectories/test_response_item_token_id.py diff --git a/docs/features/additional-histories.mdx b/docs/features/additional-histories.mdx index 86c0f583d..15d1b16f7 100644 --- a/docs/features/additional-histories.mdx +++ b/docs/features/additional-histories.mdx @@ -196,7 +196,8 @@ path, but cannot claim a complete rendered role mask. Recovering those roles requires the original request to reproduce the complete recorded prompt, with only independently proved whitespace retokenization permitted. Matching an assistant body alone does not prove an altered role header. If that proof fails, -ART refuses to attach recorded logprobs to a changed prompt. +ART refuses that rendered role certification; this check does not require +reconstructing text on the complete offline native path. Custom STOP encoders and terminator decoders must not change already-consumed source tokens, logprobs, model or stop evidence. ART checks each source around its @@ -204,11 +205,15 @@ STOP callback and checks all consumed evidence before returning, including logprobs used only by the rendered fallback. Multi-history tokenization also checks completed histories after later renderer callbacks. Evidence is bound when first read, including a later Responses generation read -to prove a copied suffix, visible item logprobs, and text used for prefix repair. +to prove a copied suffix, rendered item IDs and logprobs, and text used for prefix +repair. STOP admission also retains the source keys read before its encoder. It is not a snapshot of every future response before that response is used. Ordered original requests and history context, including roles, tools, render settings and protocol selectors, are captured before callbacks and checked -before reuse. Unexpected callback exceptions keep their original identity. +before reuse. Rendering checks its current working projection and the actual +callback arguments before and immediately after callbacks, including optional +probes; a later encoder cannot hide an earlier edit by restoring it. Unexpected +callback exceptions keep their original identity. These checks refuse stale results; they do not make callbacks or source objects immutable. Context snapshots reuse shared containers only within one observation; callback-free native assembly retains its bounded response-evidence reuse. diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index 7649e0c92..7c85e4f55 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -9,10 +9,15 @@ from enum import Enum from functools import lru_cache from hashlib import sha256 +from inspect import getattr_static +from io import BytesIO import json import math +from operator import attrgetter, is_ +from pickle import Pickler, PicklingError import re import threading +from types import FunctionType from typing import TYPE_CHECKING, Any, Literal, Protocol, cast import warnings @@ -803,6 +808,8 @@ def _merge_recorded_request_roles( output_mask: Sequence[bool], stop_mask: Sequence[bool], length_stop_mask: Sequence[bool], + *, + prompt: Sequence[int] | None, ) -> bool: # Sampled responses own their native flags. Request-only assistant roles # must retain a complete prefix proof, including roles between responses. @@ -813,7 +820,9 @@ def _merge_recorded_request_roles( ) if not output and not length_stop and (assistant or stop) ] - end = roles[-1][0] + 1 if roles else 0 + if roles and prompt is None: + return False + end = max(roles[-1][0] + 1, len(prompt or ())) if roles else 0 if exact.tokens[:end] != list(rendered[:end]) or any( exact.flags[index] & (TokenFlag.SAMPLED | TokenFlag.OUTPUT) for index, _ in roles @@ -997,6 +1006,133 @@ def __call__( ) -> None: ... +class _PlainRenderPickler(Pickler): + def reducer_override(self, value: object) -> Any: + # Exact builtins use the C traversal. Never call custom reducers. + raise TypeError("A typed rendering snapshot is required") + + +class _RenderContextGuard: + """Keep one working projection stable across external callbacks.""" + + @staticmethod + def plain(value: object, *, memo: bool) -> bytes: + stream = BytesIO() + pickler = _PlainRenderPickler(stream, protocol=4) + pickler.fast = not memo + pickler.dump(value) + return stream.getvalue() + + @classmethod + def snapshot(cls, value: object) -> tuple: + # Both encodings are local observations, never deserialized. The fresh + # C memo compresses repeated containers; the value-only encoding permits + # equal-valued copies with a different alias layout. + try: + return "plain", cls.plain(value, memo=True), cls.plain(value, memo=False) + except (TypeError, ValueError, RecursionError, PicklingError): + return "typed", _tokenization_context(value) + + def __init__(self, read: Callable[[], object]) -> None: + self.read = read + self.expected = self.snapshot(read()) + self.failed: BaseException | None = None + + def check_value(self, value: object, expected: tuple) -> None: + if self.failed is not None: + raise self.failed + try: + unchanged = ( + self.plain(value, memo=True) == expected[1] + or self.plain(value, memo=False) == expected[2] + if expected[0] == "plain" + else _tokenization_context(value) == expected[1] + ) + except (TypeError, ValueError, RecursionError, PicklingError): + unchanged = False + if not unchanged: + self.failed = ValueError( + "Rendering context changed during tokenization callback" + ) + raise self.failed + + def check(self) -> None: + if self.failed is not None: + raise self.failed + try: + value = self.read() + except BaseException as error: + self.failed = error + raise + self.check_value(value, self.expected) + + def reset(self) -> None: + # Only ART's intentional projection replacement may establish a new + # baseline, after checking the old projection and the rendered probe. + if self.failed is not None: + raise self.failed + self.expected = self.snapshot(self.read()) + + def call(self, function: Callable[..., Any], *args: Any, **kwargs: Any) -> Any: + self.check() + arguments = [list(args), kwargs] + expected = self.snapshot(arguments) + try: + result = function(*args, **kwargs) + except BaseException as error: + try: + self.check() + self.check_value(arguments, expected) + except BaseException: + # Optional renderer probes can catch the original exception. + # Keep its identity and prevent a later fallback blessing edits. + self.failed = error + raise + self.check() + self.check_value(arguments, expected) + return result + + +def _plain_tokenizer_attribute(tokenizer: object, name: str) -> tuple[bool, object]: + kind = type(tokenizer) + if type(kind) is not type: + return False, None + static = getattr_static(tokenizer, name, None) + pure = ( + getattr_static(kind, "__getattribute__") is object.__getattribute__ + and not any("__getattr__" in vars(parent) for parent in kind.__mro__) + and ( + type(static) is FunctionType + or ( + type(type(static)) is type + and getattr_static(type(static), "__get__", None) is None + ) + ) + ) + return pure, static + + +class _RenderingTokenizer: + def __init__(self, tokenizer: Tokenizer, guard: _RenderContextGuard) -> None: + self.tokenizer, self.guard = tokenizer, guard + + def __call__(self, *args: Any, **kwargs: Any) -> Any: + return self.guard.call(self.tokenizer, *args, **kwargs) + + def __getattr__(self, name: str) -> Any: + pure_lookup, _ = _plain_tokenizer_attribute(self.tokenizer, name) + value = ( + getattr(self.tokenizer, name) + if pure_lookup + else self.guard.call(lambda: getattr(self.tokenizer, name)) + ) + return ( + (lambda *args, **kwargs: self.guard.call(value, *args, **kwargs)) + if callable(value) + else value + ) + + @dataclass class _TraceBuilder: trace: _HistoryTokenizationTrace | None = None @@ -1090,13 +1226,17 @@ def set( *, tokenizer: Tokenizer | None = None, ) -> None: + if type(tokenizer) is _RenderingTokenizer: + tokenizer.guard.check() if self.track_sources: self.consume_sources( sources, require_supported_context=tokenizer is not None ) elif self.validate_sources is not None: self.validate_sources(None, require_supported_context=tokenizer is not None) - self.tokenizer = tokenizer + self.tokenizer = ( + tokenizer.tokenizer if type(tokenizer) is _RenderingTokenizer else tokenizer + ) trace = _HistoryTokenizationTrace(source_keys=source_keys, sources=sources) trace.validate(tokenized) self.trace = trace @@ -3643,12 +3783,20 @@ def _native_nonterminal_stops_known( def _history_needs_synthetic_stop( - history: History, tokenizer: Tokenizer | None + history: History, + tokenizer: Tokenizer | None, + *, + _trace: _TraceBuilder | None = None, ) -> bool: - if ( - tokenizer is None - or not callable(getattr(tokenizer, "apply_chat_template", None)) - or not _terminator_ids(tokenizer) + if tokenizer is None or not callable( + getattr(tokenizer, "apply_chat_template", None) + ): + return False + ledger = _trace or _TraceBuilder(track_sources=False) + if not ( + ledger.checked(_terminator_ids, tokenizer) + if _stop_uses_callback(None, tokenizer) + else _terminator_ids(tokenizer) ): return False sources: Sequence[object] @@ -3659,6 +3807,13 @@ def _history_needs_synthetic_stop( else: return False seen: set[_SampledSourceKey] = set() + pending: list[tuple[_SampledSourceKey, object]] = [] + + def flush() -> None: + if pending: + ledger.consume_sources(pending) + pending.clear() + for source in sources: if source is None or not _source_is_sampled(source): continue @@ -3666,16 +3821,32 @@ def _history_needs_synthetic_stop( if source_key in seen: continue seen.add(source_key) + pending.append((source_key, source)) if _source_stop_evidence(source, source_key)[0] != "stop": continue output = _source_output_tokens(source, source_key) - if output is not None and not _sampled_stop_suffix( + if output is None: + continue + callback = _stop_uses_callback( + _source_stop_evidence(source, source_key)[1], tokenizer + ) + if callback: + flush() + count = ( + ledger.checked + if callback + else lambda function, *args, **kwargs: function(*args, **kwargs) + )( + _sampled_stop_suffix, output, source=source, source_key=source_key, tokenizer=tokenizer, - ): + ) + if not count: + flush() return True + flush() return False @@ -4556,10 +4727,13 @@ def _tokenization_context( def snapshot(item: object) -> object: kind = type(item) - if kind in (str, int, bool, bytes, type(None), datetime): - return kind, item - if kind is float: - return kind, repr(item) + previous = observed.get(id(item)) + if previous is not None and previous[0] is item: + return previous[1] + if kind in (str, int, bool, bytes, type(None), datetime, float): + result = kind, repr(item) if kind is float else item + observed[id(item)] = item, result + return result if isinstance(item, Enum): if getattr(item, "__objclass__", kind) is not kind: raise TypeError("Unsupported enum tokenization context") @@ -4580,9 +4754,6 @@ def snapshot(item: object) -> object: else: scalar = bytes.__bytes__(item) return kind, scalar, snapshot(getattr(item, "__dict__", None)) - previous = observed.get(id(item)) - if previous is not None and previous[0] is item: - return previous[1] if kind in (list, tuple): result = kind, tuple(snapshot(child) for child in cast(Sequence, item)) elif isinstance(item, Mapping): @@ -4630,6 +4801,20 @@ def validate(require_supported: bool) -> None: return validate +def _rendered_response_evidence(source: object) -> object: + exchange = _source_exchange(source) + if not isinstance(exchange, ResponsesExchange): + return None + selected = _responses_source_outputs(source) + # Item-local IDs, bytes, text and logprobs can all participate in rendered + # evidence even when the native generation has its own aggregate record. + return ( + [exchange.response.output[index] for index in selected[1]] + if selected is not None + else exchange.response.output + ) + + def _sampled_source_validator( sources: Mapping[_SampledSourceKey, object] | Sequence[tuple[_SampledSourceKey, object]], @@ -4652,6 +4837,14 @@ def _sampled_source_validator( request = _tokenization_context(exchange.request, _observed=observed) except TypeError: request = None + try: + plain_request = ( + _RenderContextGuard.plain(exchange.request, memo=False) + if request is not None + else None + ) + except (TypeError, ValueError, RecursionError, PicklingError): + plain_request = None expected.setdefault(key, []).append( ( source, @@ -4659,6 +4852,7 @@ def _sampled_source_validator( exchange.model, _source_stop_evidence(source, key), request, + plain_request, _tokenization_context( { name: exchange.request.get(name) @@ -4667,12 +4861,7 @@ def _sampled_source_validator( ) if selected_request_fields is not None else None, - _tokenization_context( - _visible_logprobs( - exchange, - source=None if isinstance(source, Exchange) else source, - ) - ) + _tokenization_context(_rendered_response_evidence(source)) if rendered_evidence and isinstance(exchange, ResponsesExchange) else None, ) @@ -4693,6 +4882,7 @@ def validate( model, stop, request, + plain_request, selected_request, visible, ) in expected[key]: @@ -4712,16 +4902,11 @@ def validate( ) if ( visible is not None - and _tokenization_context( - _visible_logprobs( - exchange, - source=None if isinstance(source, Exchange) else source, - ) - ) + and _tokenization_context(_rendered_response_evidence(source)) != visible ): raise ValueError( - "Rendered source logprobs changed during tokenization callback" + "Rendered source evidence changed during tokenization callback" ) if request is None: if require_supported_context: @@ -4730,6 +4915,17 @@ def validate( ) continue current_request = exchange.request + if selected is None and plain_request is not None: + try: + # Equality only: the original typed snapshot established + # admission. Any changed/rich shape uses that same proof. + if ( + _RenderContextGuard.plain(current_request, memo=False) + == plain_request + ): + continue + except (TypeError, ValueError, RecursionError, PicklingError): + pass if selected is not None and selected_request_fields is not None: current_request = { name: exchange.request.get(name) @@ -5733,16 +5929,44 @@ def _tokenize_chat_view( # Keep every consumed source bound even for a standalone/single history. _trace.track_sources = True _validate_history_sources(history) + ledger = _trace or _TraceBuilder(track_sources=False) + _trace = ledger + if ledger.validate_context is None: + ledger.validate_context = _tokenization_context_validator( + [history, chat_template, chat_template_kwargs] + ) + # Current projected sources have already been inspected for admission. + consumed_keys = { + id(source): _sampled_source_key(source) + for source in history.message_sources + if source is not None and _source_is_sampled(source) + } + consumed_sources = [ + (consumed_keys[id(source)], source) + for source in history.message_sources + if id(source) in consumed_keys + ] + ledger.consume_sources( + consumed_sources, + selected_request_fields=("tools", "chat_template", "chat_template_kwargs"), + require_supported_context=False, + # Only Responses has additional item-local rendered evidence to admit + # after native-boundary fallback. Other protocols are already complete. + rendered_evidence=not any( + isinstance(_source_exchange(source), ResponsesExchange) + for _, source in consumed_sources + ), + ) config = ( _TokenizerConfig(base_model or history.model or "") if tokenizer is not None or (base_model is not None and chat_template is not None) - else _tokenizer_config(history.model or "", base_model) + else ledger.checked(_tokenizer_config, history.model or "", base_model) ) if tokenizer is None: if not history.model and base_model is None: raise ValueError("History tokenization requires a model or base_model") - tokenizer = _load_tokenizer(config) + tokenizer = ledger.checked(_load_tokenizer, config) assert tokenizer is not None resolved_tokenizer = tokenizer messages = [dict(message) for message in history.messages] @@ -5760,13 +5984,32 @@ def _tokenize_chat_view( explicit_kwargs.setdefault("thinking_budget", budget) template = chat_template or history.chat_template or config.chat_template if template is None: - tokenizer_template = getattr(resolved_tokenizer, "chat_template", None) + pure_template, tokenizer_template = _plain_tokenizer_attribute( + resolved_tokenizer, "chat_template" + ) + if not pure_template: + tokenizer_template = ledger.checked( + getattr, resolved_tokenizer, "chat_template", None + ) if isinstance(tokenizer_template, str): template = tokenizer_template original_template = template - template, normalization_template, defaults = _resolved_chat_template( - resolved_tokenizer, template, history.tools + pure_template, tokenizer_template = _plain_tokenizer_attribute( + resolved_tokenizer, "chat_template" ) + if ( + pure_template + and type(tokenizer_template) in (str, type(None)) + and type(template) in (str, type(None)) + ): + # Plain unnamed template normalization has no tokenizer callback. + template, normalization_template, defaults = _resolved_chat_template( + resolved_tokenizer, template, history.tools + ) + else: + template, normalization_template, defaults = ledger.checked( + _resolved_chat_template, resolved_tokenizer, template, history.tools + ) kwargs = { **defaults, **explicit_kwargs, @@ -5774,6 +6017,48 @@ def _tokenize_chat_view( ends_with_assistant = bool(messages) and messages[-1].get("role") == "assistant" segmented = False + bound_sources = tuple(history.message_sources) + present_sources = tuple(source for source in bound_sources if source is not None) + exchange_of = attrgetter("exchange") + bound_exchanges = tuple(map(exchange_of, present_sources)) + selectors = attrgetter( + "request_index", "choice_index", "output_indices", "generation_index" + ) + + def rendering_context() -> object: + # Current message projections and source selectors are linear in the + # view. Original growing requests are checked at their consumption sites. + if ( + len(history.message_sources) != len(bound_sources) + or not all(map(is_, history.message_sources, bound_sources)) + or not all(map(is_, map(exchange_of, present_sources), bound_exchanges)) + ): + raise ValueError("Rendering context changed during tokenization callback") + return [ + messages, + history.messages, + history.model, + history.tools, + history.chat_template, + history.chat_template_kwargs, + template, + kwargs, + tuple(map(selectors, present_sources)), + ] + + rendering_guard = _RenderContextGuard(rendering_context) + original_context_validator = ledger.validate_context + + def validate_rendering_context(active: bool) -> None: + assert original_context_validator is not None + original_context_validator(active) + rendering_guard.check() + + ledger.validate_context = validate_rendering_context + resolved_tokenizer = cast( + Tokenizer, _RenderingTokenizer(resolved_tokenizer, rendering_guard) + ) + def raw_render( selected_messages: list[dict[str, Any]], *, add_generation_prompt: bool ) -> list[int]: @@ -5827,32 +6112,14 @@ def render_text( ): return recorded - # Generic rendering consumes the complete projected response set. Retain - # that evidence before callbacks, rather than blessing fresh keys afterward. - consumed_keys = { - id(source): _sampled_source_key(source) - for source in history.message_sources - if source is not None and _source_is_sampled(source) - } - consumed_sources = [ - (consumed_keys[id(source)], source) - for source in history.message_sources - if id(source) in consumed_keys - ] - if _trace is not None: - _trace.consume_sources( - consumed_sources, - selected_request_fields=("tools", "chat_template", "chat_template_kwargs"), - rendered_evidence=True, - ) - assert _trace.validate_sources is not None - validate_consumed = _trace.validate_sources - else: - validate_consumed = _sampled_source_validator( + if not ledger.rendered_evidence: + ledger.consume_sources( consumed_sources, selected_request_fields=("tools", "chat_template", "chat_template_kwargs"), rendered_evidence=True, ) + assert ledger.validate_sources is not None + validate_consumed = ledger.validate_sources def consumed_source_key(source: object) -> _SampledSourceKey: key = consumed_keys[id(source)] @@ -5873,8 +6140,10 @@ def segmented_render( selected_messages: list[dict[str, Any]], *, add_generation_prompt: bool ) -> tuple[list[int], list[bool]]: try: - if cacheable_chat_template( - resolved_tokenizer, template, history.tools, kwargs, selected_messages + if rendering_guard.call( + lambda: cacheable_chat_template( + tokenizer, template, history.tools, kwargs, selected_messages + ) ): # Normalization is message-local. Only the admitted nonmutating # renderer may share its normalized messages between prefixes. @@ -6002,7 +6271,9 @@ def probe_render( add_generation_prompt=not ends_with_assistant, ) if aliased_render is not None: + rendering_guard.check() messages = aliased_messages + rendering_guard.reset() rendered = aliased_render if any( message.get("role") == "assistant" @@ -6044,7 +6315,9 @@ def probe_render( add_generation_prompt=not ends_with_assistant, ) if merged_render is not None: + rendering_guard.check() messages = merged_messages + rendering_guard.reset() rendered = merged_render part_ids_cache: dict[str, list[int]] = {} @@ -6084,14 +6357,17 @@ def part_ids(text: str) -> list[int]: prompt_cache: dict[int, list[int] | None] = {} output_cache: dict[int, tuple[list[int] | None, list[float]]] = {} - def source_prompt_tokens(source: object) -> list[int] | None: - if id(source) in consumed_keys: - consumed_source_key(source) + def cached_source_prompt(source: object) -> list[int] | None: key = id(source) if key not in prompt_cache: prompt_cache[key] = _chat_source_prompt_tokens(source) return prompt_cache[key] + def source_prompt_tokens(source: object) -> list[int] | None: + if id(source) in consumed_keys: + consumed_source_key(source) + return cached_source_prompt(source) + def source_output_tokens( source: object, ) -> tuple[list[int] | None, list[float]]: @@ -6102,7 +6378,7 @@ def source_output_tokens( output_cache[key] = _chat_source_full_tokens(source) return output_cache[key] - def source_matches_context(source: object) -> bool: + def current_source_context(source: object) -> bool: exchange = getattr(source, "exchange", None) if not isinstance( exchange, (ChatCompletionsExchange, MessagesExchange, ResponsesExchange) @@ -6129,6 +6405,16 @@ def source_matches_context(source: object) -> bool: ) ) + def source_matches_context(source: object) -> bool: + consumed_source_key(source) + return current_source_context(source) + + def matching_source_prompt(source: object) -> list[int] | None: + # One selected-source check covers this adjacent setting/cache read; + # there is no renderer or tokenizer callback between the two values. + consumed_source_key(source) + return cached_source_prompt(source) if current_source_context(source) else None + canonical_rendered = rendered exact_prefix_length = 0 canonical_prefix_length = 0 @@ -6185,6 +6471,7 @@ def source_matches_context(source: object) -> bool: ResponsesExchange, ), ) + validate_consumed(None) request_messages, request_tools = _request_messages( exchange ) @@ -6931,6 +7218,9 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: output_mask, stop_mask, length_stop_mask, + prompt=source_prompt_tokens( + history.message_sources[sampled_message_indices[-1]] + ), ): return exact if ( @@ -6953,6 +7243,7 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: validate_context = _tokenization_context_validator(history) masks = None if prompt and exact.tokens[: len(prompt)] == prompt: + validate_consumed(None) request_messages, request_tools = _request_messages(exchange) try: masks = _recorded_prompt_role_masks( @@ -7350,7 +7641,7 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: and isinstance(source_exchange, ChatCompletionsExchange) and rendered[start:end] != full_exact ): - _warn_prefix_retokenization() + rendering_guard.call(_warn_prefix_retokenization) search_cursor = end continue @@ -7455,7 +7746,7 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: (ChatCompletionsExchange, ResponsesExchange, MessagesExchange), ): evidence = _visible_token_evidence( - tokenizer, + resolved_tokenizer, exchange, source=source, sampled_text=text, @@ -7465,7 +7756,7 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: else: logprobs = ( _align_visible_logprobs( - tokenizer, + resolved_tokenizer, replacement, exchange, source=source, @@ -7665,13 +7956,9 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: source_keys.extend([None] * (len(rendered) - cursor)) exact_coverage_length = 0 for source in history.message_sources: - if ( - source is None - or not _source_is_sampled(source) - or not source_matches_context(source) - ): + if source is None or not _source_is_sampled(source): continue - source_prompt = source_prompt_tokens(source) + source_prompt = matching_source_prompt(source) if ( source_prompt is not None and token_ids[: len(source_prompt)] == source_prompt @@ -8146,7 +8433,9 @@ def _tokenize_history( can_render = tokenizer is None or callable( getattr(tokenizer, "apply_chat_template", None) ) - needs_synthetic_stop = _history_needs_synthetic_stop(history, tokenizer) + needs_synthetic_stop = _history_needs_synthetic_stop( + history, tokenizer, _trace=_trace + ) 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. diff --git a/tests/unit/trajectories/test_interior_request_prompt.py b/tests/unit/trajectories/test_interior_request_prompt.py new file mode 100644 index 000000000..f7241f310 --- /dev/null +++ b/tests/unit/trajectories/test_interior_request_prompt.py @@ -0,0 +1,78 @@ +from copy import deepcopy +import math +from typing import cast + +from openai.types.chat import ChatCompletionMessageParam +import pytest +from test_tokenize import _CharacterTemplateTokenizer, _chat_exchange + +import art.trajectories as tr + + +@pytest.mark.parametrize("change", ["none", "after_role", "inside_role"]) +def test_interior_roles_require_the_complete_final_request(change): + tokenizer = _CharacterTemplateTokenizer() + first_messages = [{"role": "user", "content": "q0"}] + first_prompt = tokenizer.apply_chat_template( + first_messages, add_generation_prompt=True + ) + assert isinstance(first_prompt, list) + output = tokenizer._encode("answer§") + first = _chat_exchange(first_prompt, output) + first.request["messages"] = cast( + list[ChatCompletionMessageParam], deepcopy(first_messages) + ) + messages = [ + *first_messages, + {"role": "assistant", "content": "answer"}, + {"role": "user", "content": "u1"}, + {"role": "assistant", "content": "history"}, + {"role": "user", "content": "last"}, + ] + recorded = deepcopy(messages) + if change == "after_role": + recorded[-1]["content"] = "LAST" + elif change == "inside_role": + recorded[-2]["content"] = "HISTORY" + prompt = tokenizer.apply_chat_template(recorded, add_generation_prompt=True) + assert isinstance(prompt, list) + second = _chat_exchange(prompt, output, offset=1) + second.request["messages"] = cast( + list[ChatCompletionMessageParam], deepcopy(messages) + ) + value = tr.Trajectory( + exchanges=tr.TrajectoryExchanges(chat_completions=[first, second]) + ) + before = value.model_dump_json() + if change != "none": + with pytest.raises(ValueError, match="[Rr]equest roles|assistant boundaries"): + value.tokenize(tokenizer=tokenizer, multi_history=True) + else: + histories = value.tokenize(tokenizer=tokenizer, multi_history=True).histories + assert len(histories) == 1 + actual = histories[0] + assert actual.tokens == prompt + output + sampled = ( + tr.TokenFlag.EXACT + | tr.TokenFlag.SAMPLED + | tr.TokenFlag.OUTPUT + | tr.TokenFlag.ASSISTANT + ) + for start in (len(first_prompt), len(prompt)): + assert actual.logprobs[start : start + len(output)] == [ + -token / 10 for token in output + ] + assert all( + flag & sampled == sampled + for flag in actual.flags[start : start + len(output)] + ) + interior = range( + len(first_prompt) + len(output) + len("u1"), len(prompt) - len("last") + ) + assert all(actual.flags[i] & tr.TokenFlag.ASSISTANT for i in interior) + assert all( + not actual.flags[i] & (tr.TokenFlag.SAMPLED | tr.TokenFlag.OUTPUT) + for i in interior + ) + assert all(math.isnan(actual.logprobs[i]) for i in interior) + assert value.model_dump_json() == before diff --git a/tests/unit/trajectories/test_renderer_context_callbacks.py b/tests/unit/trajectories/test_renderer_context_callbacks.py new file mode 100644 index 000000000..bfa68e050 --- /dev/null +++ b/tests/unit/trajectories/test_renderer_context_callbacks.py @@ -0,0 +1,410 @@ +from __future__ import annotations + +from pickle import PickleBuffer +from typing import Any, cast + +import pytest +from test_tokenize import _chat_exchange + +import art.trajectories as tr + + +@pytest.mark.parametrize("mutate", [False, True]) +def test_render_context_mutation_cannot_be_restored_by_later_encoder(mutate): + exchange = _chat_exchange([1], [2]) + exchange.request["messages"] = [{"role": "user", "content": "q"}] + choice = exchange.response.choices[0] + assert choice.model_extra is not None + choice.model_extra.pop("prompt_token_ids") + value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(chat_completions=[exchange])) + before = value.model_dump_json() + calls = [] + + class Tokenizer: + def apply_chat_template(self, messages, **kwargs): + calls.append("render") + if mutate: + exchange.request["messages"][0]["content"] = "changed" + messages[0]["content"] = "changed" + return [3 if mutate else 1, 2] + + def __call__(self, text, **kwargs): + calls.append(text) + if mutate and text == "changed": + exchange.request["messages"][0]["content"] = "q" + return [{"q": 1, "changed": 3, "answer": 2}[text]] + + if mutate: + with pytest.raises(ValueError, match="[Cc]ontext changed"): + actual = value.tokenize( + tokenizer=cast(Any, Tokenizer()), chat_template="custom" + ) + assert actual.tokens == [3, 2] + assert actual.logprobs[-1] == -0.2 + assert actual.flags[-1] & tr.TokenFlag.SAMPLED + assert value.model_dump_json() == before + else: + actual = value.tokenize( + tokenizer=cast(Any, Tokenizer()), chat_template="custom" + ) + assert actual.tokens == [1, 2] + assert actual.logprobs[-1] == -0.2 + assert actual.flags[-1] & tr.TokenFlag.SAMPLED + assert value.model_dump_json() == before + assert "render" in calls + + +@pytest.mark.parametrize("mutate", [False, True]) +def test_optional_segment_encoder_cannot_swallow_mutated_projection(mutate): + from test_tokenize import _repeated_text_rerender_history, _RepeatedTextTokenizer + + error = ValueError("encoder capability unavailable") + calls = [] + + class Tokenizer(_RepeatedTextTokenizer): + working = None + + def apply_chat_template(self, messages, **kwargs): + if self.working is None: + self.working = messages + return super().apply_chat_template(messages, **kwargs) + + def __call__(self, text, **kwargs): + if "" in text and not kwargs.get("return_offsets_mapping"): + calls.append(text) + if mutate: + assert self.working is not None + self.working[0]["content"] = "changed" + raise error + return super().__call__(text, **kwargs) + + history = _repeated_text_rerender_history(2) + if mutate: + with pytest.raises(ValueError) as raised: + history.tokenize(tokenizer=Tokenizer()) + assert raised.value is error + else: + actual = history.tokenize(tokenizer=Tokenizer()) + assert actual.tokens and any( + flag & tr.TokenFlag.SAMPLED for flag in actual.flags + ) + assert calls + + +@pytest.mark.parametrize("api", ["history", "trajectory", "trace"]) +@pytest.mark.parametrize("mutate", [False, True]) +def test_stop_decision_binds_logprobs_before_encoding(api, mutate): + from test_tokenize import _CharacterTemplateTokenizer + + from art.trajectories import _tokenize as module + + calls = [] + exchange = _chat_exchange([1], [2, 9]) + choice = exchange.response.choices[0] + assert choice.model_extra is not None + assert choice.logprobs is not None and choice.logprobs.content + choice.model_extra["stop_reason"] = "§" + rows = choice.logprobs.content + original = rows[0].logprob + + class Tokenizer(_CharacterTemplateTokenizer): + def __call__(self, text, **kwargs): + if text == "§": + calls.append(True) + if mutate and len(calls) == 1: + rows[0].logprob = -99.0 + return super().__call__(text, **kwargs) + + value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(chat_completions=[exchange])) + tokenizer = Tokenizer() + + def invoke(): + if api == "history": + return value.chat_completions_history().tokenize(tokenizer=tokenizer) + if api == "trace": + return module._tokenize_trajectory_with_trace(value, tokenizer=tokenizer)[0] + return value.tokenize(multi_history=True, tokenizer=tokenizer) + + if mutate: + with pytest.raises(ValueError, match="[Ss]ampled source changed"): + invoke() + assert rows[0].logprob == -99.0 + else: + actual = invoke() + history = actual if api == "history" else actual.histories[0] + assert history.tokens == [1, 2, 9] + assert history.logprobs[1] == original + assert calls + + +@pytest.mark.parametrize( + ("before", "after"), + [ + (1, True), + (1, 1.0), + (-0.0, 0.0), + ("😀", "\ud83d\ude00"), + (b"a", bytearray(b"a")), + (b"a", PickleBuffer(b"a")), + ({"a": 1, "b": 2}, {"b": 2, "a": 1}), + ], +) +@pytest.mark.parametrize("mutate", [False, True]) +def test_renderer_context_keeps_types_order_and_unicode(before, after, mutate): + options = {"value": before} + exchange = _chat_exchange([1], [2]) + exchange.request["messages"] = [{"role": "user", "content": "q"}] + value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(chat_completions=[exchange])) + calls = [] + + class Tokenizer: + def apply_chat_template(self, messages, **kwargs): + calls.append(True) + assert kwargs["options"] is options + if mutate: + options["value"] = after + return [3 if mutate else 1, 2] + + def __call__(self, text, **kwargs): + # A later callback must not hide a changed rendering argument. + if mutate: + options["value"] = before + return [2 if text == "answer" else 3 if mutate else 1] + + if mutate: + with pytest.raises(ValueError, match="[Cc]ontext changed"): + value.tokenize( + tokenizer=cast(Any, Tokenizer()), + chat_template="custom", + chat_template_kwargs={"options": options}, + ) + else: + actual = value.tokenize( + tokenizer=cast(Any, Tokenizer()), + chat_template="custom", + chat_template_kwargs={"options": options}, + ) + assert actual.tokens == [1, 2] and actual.logprobs[-1] == -0.2 + assert calls + + +@pytest.mark.parametrize("lookup", ["property", "custom"]) +@pytest.mark.parametrize("mutate", [False, True]) +def test_tokenizer_attribute_callbacks_keep_projection_guard(lookup, mutate): + from test_tokenize import _repeated_text_rerender_history, _RepeatedTextTokenizer + + calls = [] + + class Tokenizer(_RepeatedTextTokenizer): + working = None + + def apply_chat_template(self, messages, **kwargs): + if self.working is None: + self.working = messages + return super().apply_chat_template(messages, **kwargs) + + def observe(self): + if self.working is not None: + calls.append(True) + if mutate: + self.working[0]["content"] = "changed" + return None + + if lookup == "property": + setattr(Tokenizer, "eos_token_id", property(lambda self: self.observe())) + else: + + def read(self, name): + if name == "eos_token_id": + return self.observe() + return object.__getattribute__(self, name) + + Tokenizer.__getattribute__ = read + history = _repeated_text_rerender_history(2) + if mutate: + with pytest.raises(ValueError, match="[Cc]ontext changed"): + history.tokenize(tokenizer=Tokenizer()) + else: + actual = history.tokenize(tokenizer=Tokenizer()) + assert actual.tokens and any( + flag & tr.TokenFlag.SAMPLED for flag in actual.flags + ) + assert calls + + +@pytest.mark.parametrize("shape", ["subclass", "cycle"]) +@pytest.mark.parametrize("mutate", [False, True]) +def test_render_context_rich_fallback_never_invokes_reducers(shape, mutate): + class Text(str): + state: int + + def __reduce_ex__(self, protocol): + raise AssertionError("render guard must not call custom reducers") + + text = Text("stable") + text.state = 1 + options: dict[str, Any] = {"value": text if shape == "subclass" else []} + exchange = _chat_exchange([1], [2]) + exchange.request["messages"] = [{"role": "user", "content": "q"}] + value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(chat_completions=[exchange])) + calls = [] + + class Tokenizer: + def apply_chat_template(self, messages, **kwargs): + calls.append(True) + if mutate: + if shape == "subclass": + text.state = 2 + else: + options["value"].append(options["value"]) + return [3 if mutate else 1, 2] + + def __call__(self, encoded, **kwargs): + text.state = 1 + if shape == "cycle": + options["value"].clear() + return [2 if encoded == "answer" else 3 if mutate else 1] + + if mutate: + with pytest.raises(ValueError, match="[Cc]ontext changed"): + value.tokenize( + tokenizer=cast(Any, Tokenizer()), + chat_template="custom", + chat_template_kwargs={"options": options}, + ) + else: + actual = value.tokenize( + tokenizer=cast(Any, Tokenizer()), + chat_template="custom", + chat_template_kwargs={"options": options}, + ) + assert actual.tokens == [1, 2] and actual.logprobs[-1] == -0.2 + assert calls + + +@pytest.mark.parametrize("mutate", [False, True]) +def test_render_guard_accepts_equal_copies_but_checks_their_values(mutate): + shared = {"value": 1} + options = [shared, shared] + exchange = _chat_exchange([1], [2]) + exchange.request["messages"] = [{"role": "user", "content": "q"}] + value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(chat_completions=[exchange])) + calls = [] + + class Tokenizer: + def apply_chat_template(self, messages, **kwargs): + calls.append(True) + options[1] = {"value": 2 if mutate else 1} + return [3 if mutate else 1, 2] + + def __call__(self, text, **kwargs): + options[1] = shared + return [2 if text == "answer" else 3 if mutate else 1] + + if mutate: + with pytest.raises(ValueError, match="[Cc]ontext changed"): + value.tokenize( + tokenizer=cast(Any, Tokenizer()), + chat_template="custom", + chat_template_kwargs={"options": options}, + ) + else: + actual = value.tokenize( + tokenizer=cast(Any, Tokenizer()), + chat_template="custom", + chat_template_kwargs={"options": options}, + ) + assert actual.tokens == [1, 2] and actual.logprobs[-1] == -0.2 + assert calls + + +@pytest.mark.parametrize("mutate", [False, True]) +def test_visible_evidence_encoder_cannot_restore_working_projection(mutate): + from openai.types.responses import Response + from test_tokenize import _multi_output_responses_chat_history + + projected = _multi_output_responses_chat_history() + source = next(source for source in projected.message_sources if source is not None) + exchange = source.exchange + assert isinstance(exchange, tr.ResponsesExchange) + data = exchange.response.model_dump(mode="python") + for item, text, lp in zip( + data["output"], ["first", "second"], [-0.1, -0.2], strict=True + ): + item["content"][0]["logprobs"] = [ + { + "token": text, + "logprob": lp, + "bytes": list(text.encode()), + "top_logprobs": [], + } + ] + exchange.response = Response.model_validate(data) + value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(responses=[exchange])) + error = ValueError("offsets unavailable") + calls = [] + + class Tokenizer: + working = None + + def apply_chat_template(self, messages, **kwargs): + if self.working is None: + self.working = messages + return [token for message in messages for token in self(message["content"])] + + def __call__(self, text, **kwargs): + token = {"turn 0": 1, "first": 20, "second": 30}[text] + if kwargs.get("return_offsets_mapping"): + calls.append(text) + if mutate and text == "first": + assert self.working is not None + self.working[0]["content"] = "changed" + raise error + return {"input_ids": [token], "offset_mapping": [(0, len(text))]} + if self.working is not None: + self.working[0]["content"] = "turn 0" + return [token] + + if mutate: + with pytest.raises(ValueError) as raised: + value.tokenize(tokenizer=Tokenizer(), chat_template="custom") + assert raised.value is error + else: + actual = value.tokenize(tokenizer=Tokenizer(), chat_template="custom") + assert actual.tokens == [1, 20, 30] + assert actual.logprobs[1:] == [-0.1, -0.2] + assert "first" in calls + + +def test_render_guard_does_not_certify_unsupported_callback_arguments(): + from art.trajectories._tokenize import _RenderContextGuard + + class Opaque: + __slots__ = () + + calls = [] + guard = _RenderContextGuard(lambda: {"stable": True}) + with pytest.raises(TypeError, match="Unsupported mutable tokenization context"): + guard.call(lambda value: calls.append(value), Opaque()) + assert not calls + + +def test_trace_unwrap_does_not_probe_custom_tokenizer_class(): + from art.trajectories._tokenize import _TraceBuilder, _sampled_source_key + + class Tokenizer: + @property + def __class__(self): + raise AssertionError("wrapper admission must not call a custom attribute") + + exchange = _chat_exchange([1], [2]) + value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(chat_completions=[exchange])) + tokenized = value.tokenize() + builder = _TraceBuilder(track_sources=False) + tokenizer = cast(Any, Tokenizer()) + source = value.chat_completions_history().message_sources[-1] + assert source is not None + key = _sampled_source_key(source) + keys = [key if flag & tr.TokenFlag.SAMPLED else None for flag in tokenized.flags] + builder.set(tokenized, keys, {key: source}, tokenizer=tokenizer) + assert builder.tokenizer is tokenizer diff --git a/tests/unit/trajectories/test_response_item_token_id.py b/tests/unit/trajectories/test_response_item_token_id.py new file mode 100644 index 000000000..cad7d49ad --- /dev/null +++ b/tests/unit/trajectories/test_response_item_token_id.py @@ -0,0 +1,112 @@ +"""Public explicit-rendered Responses item-ID consumption regression.""" + +from __future__ import annotations + +import math +from typing import Any + +from openai.types.responses import Response, ResponseOutputMessage, ResponseOutputText +import pytest +from test_tokenize import _multi_output_responses_chat_history + +import art.trajectories as tr + + +@pytest.mark.parametrize("route", ["history", "trajectory"]) +@pytest.mark.parametrize("mutate", [False, True]) +def test_multi_output_generation_keeps_consumed_item_token_id(route, mutate): + projected = _multi_output_responses_chat_history() + source = next( + source + for source in projected.message_sources + if source is not None and source.output_indices == (0,) + ) + exchange = source.exchange + assert isinstance(exchange, tr.ResponsesExchange) + data = exchange.response.model_dump(mode="python") + # Retain one exact generation spanning both outputs: neither individual + # rendered message owns all of its native output evidence. + assert data["token_generations"][0]["output_indices"] == [0, 1] + for output, text, lp in zip( + data["output"], ["first", "second"], [-0.1, -0.2], strict=True + ): + output["content"][0]["logprobs"] = [ + { + "token": text, + "logprob": lp, + "bytes": list(text.encode()), + "top_logprobs": [], + } + ] + data["output"][0]["content"][0]["logprobs"][0]["token_id"] = 20 + exchange.response = Response.model_validate(data) + value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(responses=[exchange])) + item = exchange.response.output[0] + assert isinstance(item, ResponseOutputMessage) + content = item.content[0] + assert isinstance(content, ResponseOutputText) and content.logprobs + first_lp = content.logprobs[0] + assert first_lp.model_extra is not None + extras = first_lp.model_extra + calls = [] + + class Tokenizer: + def __call__(self, text: str, **kwargs: Any): + token = {"turn 0": 1, "first": 20, "second": 30}[text] + if kwargs.get("return_offsets_mapping"): + calls.append(text) + if mutate and text == "second": + assert first_lp.model_extra is not None + extras["token_id"] = 99 + return {"input_ids": [token], "offset_mapping": [(0, len(text))]} + return [token] + + def apply_chat_template(self, messages, **kwargs): + return [token for message in messages for token in self(message["content"])] + + def tokenize(): + if route == "history": + return value.responses_history().tokenize( + tokenizer=Tokenizer(), chat_template="custom" + ) + return value.tokenize(tokenizer=Tokenizer(), chat_template="custom") + + if mutate: + with pytest.raises( + ValueError, + match="[Ss]ampled source changed|[Cc]onsumed|[Cc]ontext changed|Rendered source evidence changed", + ): + actual = tokenize() + # If the guard misses the edit, prove this is the consumed old + # value being returned, not a benign pre-consumption refresh. + assert extras["token_id"] == 99 + assert first_lp.logprob == -0.1 + assert actual.tokens == [1, 20, 30] + assert math.isnan(actual.logprobs[0]) + assert actual.logprobs[1:] == [-0.1, -0.2] + assert actual.flags == [ + tr.TokenFlag(0), + tr.TokenFlag.ASSISTANT + | tr.TokenFlag.OUTPUT + | tr.TokenFlag.EXACT + | tr.TokenFlag.SAMPLED, + tr.TokenFlag.ASSISTANT | tr.TokenFlag.OUTPUT, + ] + assert extras["token_id"] == 99 + assert first_lp.logprob == -0.1 + else: + actual = tokenize() + assert actual.tokens == [1, 20, 30] + assert math.isnan(actual.logprobs[0]) + assert actual.logprobs[1:] == [-0.1, -0.2] + assert actual.flags == [ + tr.TokenFlag(0), + tr.TokenFlag.ASSISTANT + | tr.TokenFlag.OUTPUT + | tr.TokenFlag.EXACT + | tr.TokenFlag.SAMPLED, + tr.TokenFlag.ASSISTANT | tr.TokenFlag.OUTPUT, + ] + assert first_lp.logprob == -0.1 + assert extras["token_id"] == 20 + assert "second" in calls diff --git a/tests/unit/trajectories/test_visible_response_logprobs.py b/tests/unit/trajectories/test_visible_response_logprobs.py index e073256e6..58406d7c6 100644 --- a/tests/unit/trajectories/test_visible_response_logprobs.py +++ b/tests/unit/trajectories/test_visible_response_logprobs.py @@ -70,7 +70,7 @@ def tokenize(): if mutate: with pytest.raises( ValueError, - match="[Ss]ampled source changed|[Cc]onsumed|[Cc]ontext changed|Rendered source logprobs changed", + match="[Ss]ampled source changed|[Cc]onsumed|[Cc]ontext changed|Rendered source evidence changed", ): actual = tokenize() # If the guard misses the edit, prove this is the consumed old From da6c8665cf91808000f4a0b385382a0c5915dd3d Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 22:47:18 +0000 Subject: [PATCH 27/67] Keep loader mutation tests aligned with consumed evidence lifetime --- .../unit/trajectories/test_evidence_reuse.py | 33 ++++++++++++------- 1 file changed, 22 insertions(+), 11 deletions(-) diff --git a/tests/unit/trajectories/test_evidence_reuse.py b/tests/unit/trajectories/test_evidence_reuse.py index c30e2d597..64020ab83 100644 --- a/tests/unit/trajectories/test_evidence_reuse.py +++ b/tests/unit/trajectories/test_evidence_reuse.py @@ -102,7 +102,7 @@ def __call__(self, text, **kwargs): @pytest.mark.parametrize("override", [False, True]) -def test_render_fallback_does_not_receive_decision_evidence(monkeypatch, override): +def test_render_fallback_binds_evidence_before_loader_callback(monkeypatch, override): from test_tokenize import _character_template_history monkeypatch.setattr(module, "_WARNED_PREFIX_RETOKENIZATION", False) @@ -125,16 +125,25 @@ def load(config): "_tokenizer_config", lambda *args: module._TokenizerConfig("public/base"), ) + + def invoke(trace): + return 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, + ) + + # The loader cannot replace evidence already inspected for this call. + with pytest.raises(ValueError, match="Sampled source changed"): + invoke(module._TraceBuilder()) + # A subsequent invocation binds the edited value freshly; no decision memo + # or failed-call validator leaks into that independent tokenization. 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, - ) + result = invoke(trace) assert trace.trace is not None new_key = module._sampled_source_key(first_source) assert new_key != old_key @@ -471,7 +480,9 @@ def __call__(self, text, **kwargs): exchange.request["model"] = "changed/model" return {"input_ids": [3]} - with pytest.raises(ValueError, match="model no longer matches"): + with pytest.raises( + ValueError, match="model no longer matches|Sampled source changed" + ): history.tokenize(tokenizer=Tokenizer()) From 6864c964346e6f6e28233573b99d88e02f66eacf Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 22:54:27 +0000 Subject: [PATCH 28/67] Prove request roles without reconstructing later native bodies --- docs/features/additional-histories.mdx | 3 + src/art/trajectories/_tokenize.py | 29 +++++--- .../test_interior_request_prompt.py | 72 +++++++++++++++++++ 3 files changed, 96 insertions(+), 8 deletions(-) diff --git a/docs/features/additional-histories.mdx b/docs/features/additional-histories.mdx index 15d1b16f7..e77562a15 100644 --- a/docs/features/additional-histories.mdx +++ b/docs/features/additional-histories.mdx @@ -188,6 +188,9 @@ 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. +The role proof covers the complete recorded request containing the last +request-only assistant, including its trailing query. Later sampled native +bodies remain authoritative even when their rendered projections differ. With a supplied renderer, request-owned assistant roles are proved throughout the native stream, including between sampled responses. This requirement applies even when template normalization makes no change. A request whose text diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index 7c85e4f55..593208981 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -809,10 +809,10 @@ def _merge_recorded_request_roles( stop_mask: Sequence[bool], length_stop_mask: Sequence[bool], *, - prompt: Sequence[int] | None, + prompts: Iterable[Sequence[int] | None], ) -> bool: - # Sampled responses own their native flags. Request-only assistant roles - # must retain a complete prefix proof, including roles between responses. + # Prove the complete recorded request owning the last request-only role. + # Later sampled native bodies need not match their rendered projection. roles = [ (index, _rendered_flag(assistant, False, stop)) for index, (assistant, output, stop, length_stop) in enumerate( @@ -820,9 +820,21 @@ def _merge_recorded_request_roles( ) if not output and not length_stop and (assistant or stop) ] - if roles and prompt is None: - return False - end = max(roles[-1][0] + 1, len(prompt or ())) if roles else 0 + end = 0 + if roles: + prompt = next( + ( + prompt + for prompt in prompts + if prompt is not None + and len(prompt) > roles[-1][0] + and exact.tokens[: len(prompt)] == list(prompt) + ), + None, + ) + if prompt is None: + return False + end = len(prompt) if exact.tokens[:end] != list(rendered[:end]) or any( exact.flags[index] & (TokenFlag.SAMPLED | TokenFlag.OUTPUT) for index, _ in roles @@ -7218,8 +7230,9 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: output_mask, stop_mask, length_stop_mask, - prompt=source_prompt_tokens( - history.message_sources[sampled_message_indices[-1]] + prompts=( + source_prompt_tokens(history.message_sources[index]) + for index in sampled_message_indices ), ): return exact diff --git a/tests/unit/trajectories/test_interior_request_prompt.py b/tests/unit/trajectories/test_interior_request_prompt.py index f7241f310..4acdd78f6 100644 --- a/tests/unit/trajectories/test_interior_request_prompt.py +++ b/tests/unit/trajectories/test_interior_request_prompt.py @@ -76,3 +76,75 @@ def test_interior_roles_require_the_complete_final_request(change): ) assert all(math.isnan(actual.logprobs[i]) for i in interior) assert value.model_dump_json() == before + + +@pytest.mark.parametrize("body", ["text", "tool", "reasoning"]) +def test_initial_request_roles_do_not_reencode_later_native_bodies(body): + from openai.types.chat import ChatCompletion + + tokenizer = _CharacterTemplateTokenizer() + messages = [ + {"role": "user", "content": "intro"}, + {"role": "assistant", "content": "history"}, + {"role": "user", "content": "turn0"}, + ] + prompt = tokenizer._encode("introhistory§turn0") + output = tokenizer._encode("different native body§") + first = _chat_exchange(prompt, output) + first.request["messages"] = cast( + list[ChatCompletionMessageParam], deepcopy(messages) + ) + payload = first.response.model_dump(mode="python") + message = payload["choices"][0]["message"] + if body == "tool": + message["content"] = None + message["tool_calls"] = [ + { + "id": "public", + "type": "function", + "function": {"name": "lookup", "arguments": '{"key": 1}'}, + } + ] + elif body == "reasoning": + message["reasoning_content"] = "structured reasoning" + first.response = ChatCompletion.model_validate(payload) + final_prompt = [*prompt, *output, *tokenizer._encode("turn1")] + second = _chat_exchange(final_prompt, tokenizer._encode("answer§"), offset=1) + second.request["messages"] = cast( + list[ChatCompletionMessageParam], + [ + *deepcopy(messages), + first.response.choices[0].message.model_dump( + mode="python", exclude_none=True + ), + {"role": "user", "content": "turn1"}, + ], + ) + value = tr.Trajectory( + exchanges=tr.TrajectoryExchanges(chat_completions=[first, second]) + ) + before = value.model_dump_json() + actual = value.tokenize(tokenizer=tokenizer, multi_history=True).histories + assert len(actual) == 1 + result = actual[0] + assert result.tokens == [*final_prompt, *tokenizer._encode("answer§")] + for start, ids in [ + (len(prompt), output), + (len(final_prompt), tokenizer._encode("answer§")), + ]: + assert result.logprobs[start : start + len(ids)] == [ + -token / 10 for token in ids + ] + assert all( + flag & tr.TokenFlag.SAMPLED + for flag in result.flags[start : start + len(ids)] + ) + historical = slice(len("intro"), len("introhistory§")) + assert all(flag & tr.TokenFlag.ASSISTANT for flag in result.flags[historical]) + assert result.flags[len("introhistory")] & tr.TokenFlag.STOP + assert not any( + flag & (tr.TokenFlag.SAMPLED | tr.TokenFlag.OUTPUT) + for flag in result.flags[historical] + ) + assert all(math.isnan(lp) for lp in result.logprobs[historical]) + assert value.model_dump_json() == before From 6d140317e7a0d0cea2ea7a58632a5a8ac11fb0d6 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 23:09:17 +0000 Subject: [PATCH 29/67] Fence dynamic STOP metadata and avoid unused template normalization --- docs/features/additional-histories.mdx | 12 ++- src/art/trajectories/_tokenize.py | 34 ++++++-- .../test_messages_disposable_projection.py | 42 ++++++++++ .../test_named_template_arguments.py | 27 +++++++ .../test_renderer_context_callbacks.py | 2 +- .../test_stop_attribute_callbacks.py | 77 +++++++++++++++++++ 6 files changed, 183 insertions(+), 11 deletions(-) create mode 100644 tests/unit/trajectories/test_messages_disposable_projection.py create mode 100644 tests/unit/trajectories/test_stop_attribute_callbacks.py diff --git a/docs/features/additional-histories.mdx b/docs/features/additional-histories.mdx index e77562a15..6bbcf71cc 100644 --- a/docs/features/additional-histories.mdx +++ b/docs/features/additional-histories.mdx @@ -202,8 +202,8 @@ assistant body alone does not prove an altered role header. If that proof fails, ART refuses that rendered role certification; this check does not require reconstructing text on the complete offline native path. -Custom STOP encoders and terminator decoders must not change already-consumed -source tokens, logprobs, model or stop evidence. ART checks each source around its +Custom STOP encoders, terminator decoders and metadata lookups 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, including logprobs used only by the rendered fallback. Multi-history tokenization also checks completed histories after later renderer callbacks. @@ -213,10 +213,16 @@ repair. STOP admission also retains the source keys read before its encoder. It is not a snapshot of every future response before that response is used. Ordered original requests and history context, including roles, tools, render settings and protocol selectors, are captured before callbacks and checked -before reuse. Rendering checks its current working projection and the actual +before reuse. Chat role rendering checks its working projection and actual callback arguments before and immediately after callbacks, including optional probes; a later encoder cannot hide an earlier edit by restoring it. Unexpected callback exceptions keep their original identity. + +When recorded prompt IDs are unavailable, a supplied renderer defines the prompt +tokens. ART cannot certify that arbitrary returned IDs implement a particular +text transformation. A renderer changing a disposable, single-use message copy +does not itself invalidate that authority; changes to borrowed sources or +semantic inputs that ART reuses still require refusal. These checks refuse stale results; they do not make callbacks or source objects immutable. Context snapshots reuse shared containers only within one observation; callback-free native assembly retains its bounded response-evidence reuse. diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index 593208981..86ff0415b 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -2523,11 +2523,15 @@ def _resolved_chat_template( ) -> tuple[object, 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) + deferred = isinstance(template, dict) + configured = ( + template if deferred else 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): + deferred = False selected = select( chat_template=template if isinstance(template, str) else None, tools=tools, @@ -2547,6 +2551,8 @@ def _resolved_chat_template( "The normalized chat template is also a template name; " "cannot preserve the selected renderer without ambiguity" ) + if deferred: + configured = chat_template_with_preserved_thinking(configured) return configured, configured, defaults @@ -4956,12 +4962,26 @@ def 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)) - ) + if tokenizer is None or isinstance(reason, int) and not isinstance(reason, bool): + return False + if isinstance(reason, str) and reason: + return True + # Admission must not itself invoke a descriptor or custom lookup. Only + # absent metadata and plain scalar EOS IDs prove a callback-free lookup. + for name in ( + "eos_token_id", + "eot_token_id", + "special_tokens_map", + "convert_tokens_to_ids", + ): + pure, value = _plain_tokenizer_attribute(tokenizer, name) + if ( + not pure + or value is not None + and not (name in ("eos_token_id", "eot_token_id") and type(value) is int) + ): + return True + return False def _mark_sampled_stops( diff --git a/tests/unit/trajectories/test_messages_disposable_projection.py b/tests/unit/trajectories/test_messages_disposable_projection.py new file mode 100644 index 000000000..ece9f9d48 --- /dev/null +++ b/tests/unit/trajectories/test_messages_disposable_projection.py @@ -0,0 +1,42 @@ +from typing import Any, cast + +import pytest +from test_tokenize import _message_exchange + +import art.trajectories as tr + + +@pytest.mark.parametrize("mutate_copy", [False, True]) +def test_disposable_messages_renderer_is_prompt_authority(mutate_copy): + exchange = _message_exchange( + cast( + Any, + { + "model": "test/model", + "max_tokens": 5, + "messages": [{"role": "user", "content": "question"}], + }, + ), + token_ids=[2], + logprobs=[-0.2], + ) + original = exchange.model_dump_json() + calls = [] + + class Tokenizer: + def apply_chat_template(self, messages, **kwargs): + assert messages == [{"role": "user", "content": "question"}] + assert messages is not exchange.request["messages"] + assert messages[0] is not exchange.request["messages"][0] + if mutate_copy: + messages[0]["content"] = "changed" + calls.append(True) + return [99] + + value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(messages=[exchange])) + actual = value.tokenize(tokenizer=cast(Any, Tokenizer())) + assert calls == [True] + assert actual.tokens == [99, 2] + assert actual.logprobs[-1] == -0.2 + assert actual.flags[-1] & tr.TokenFlag.SAMPLED + assert exchange.model_dump_json() == original diff --git a/tests/unit/trajectories/test_named_template_arguments.py b/tests/unit/trajectories/test_named_template_arguments.py index 6c1e87c0d..51c0639f8 100644 --- a/tests/unit/trajectories/test_named_template_arguments.py +++ b/tests/unit/trajectories/test_named_template_arguments.py @@ -99,3 +99,30 @@ def test_selected_body_normalizes_tool_arguments_without_reselecting( settings["enable_thinking"] is True and settings["preserve_thinking"] is False for settings in tokenizer.settings ) + + +@pytest.mark.parametrize("with_tools", [False, True]) +def test_named_selection_normalizes_only_the_selected_body(monkeypatch, with_tools): + tokenizer = _NamedTemplateTokenizer() + original = deepcopy(tokenizer.chat_template) + tools = ( + [{"type": "function", "function": {"name": "lookup"}}] if with_tools else None + ) + selected = tokenizer.chat_template["tool_use" if with_tools else "default"] + normalize = _tokenize.chat_template_with_preserved_thinking + expected = normalize(selected) + calls = [] + + def observe(template): + calls.append(template) + return normalize(template) + + monkeypatch.setattr(_tokenize, "chat_template_with_preserved_thinking", observe) + selector, body, defaults = _tokenize._resolved_chat_template( + tokenizer, tokenizer.chat_template, tools + ) + assert body == expected + assert selector == (expected if expected != selected else tokenizer.chat_template) + assert defaults == {} + assert tokenizer.chat_template == original + assert calls == [selected] diff --git a/tests/unit/trajectories/test_renderer_context_callbacks.py b/tests/unit/trajectories/test_renderer_context_callbacks.py index bfa68e050..3980a952e 100644 --- a/tests/unit/trajectories/test_renderer_context_callbacks.py +++ b/tests/unit/trajectories/test_renderer_context_callbacks.py @@ -390,7 +390,7 @@ class Opaque: def test_trace_unwrap_does_not_probe_custom_tokenizer_class(): - from art.trajectories._tokenize import _TraceBuilder, _sampled_source_key + from art.trajectories._tokenize import _sampled_source_key, _TraceBuilder class Tokenizer: @property diff --git a/tests/unit/trajectories/test_stop_attribute_callbacks.py b/tests/unit/trajectories/test_stop_attribute_callbacks.py new file mode 100644 index 000000000..5691ffe0c --- /dev/null +++ b/tests/unit/trajectories/test_stop_attribute_callbacks.py @@ -0,0 +1,77 @@ +from typing import Any, cast + +import pytest +from test_tokenize import _chat_exchange + +import art.trajectories as tr + + +@pytest.mark.parametrize( + "attribute", ["eos_token_id", "eot_token_id", "special_tokens_map"] +) +@pytest.mark.parametrize("mutate", [False, True]) +def test_stop_metadata_lookup_preserves_consumed_logprobs(attribute, mutate): + exchange = _chat_exchange([1], [2, 3]) + choice = exchange.response.choices[0] + assert choice.logprobs is not None and choice.logprobs.content is not None + logprobs = choice.logprobs.content + calls = [] + + def read(_self): + calls.append(True) + if mutate: + logprobs[0].logprob = -9 + return {} if attribute == "special_tokens_map" else 3 + + tokenizer = type("Tokenizer", (), {attribute: property(read)})() + value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(chat_completions=[exchange])) + before = value.model_dump_json() + if mutate: + with pytest.raises(ValueError, match="[Ss]ampled source changed"): + actual = value.tokenize(tokenizer=cast(Any, tokenizer)) + assert actual.tokens == [1, 2, 3] + assert actual.logprobs[-2:] == [-0.2, -0.3] + assert logprobs[0].logprob == -9 + else: + actual = value.tokenize(tokenizer=cast(Any, tokenizer)) + assert actual.tokens == [1, 2, 3] + assert actual.logprobs[-2:] == [-0.2, -0.3] + assert bool(actual.flags[-1] & tr.TokenFlag.STOP) == ( + attribute != "special_tokens_map" + ) + assert value.model_dump_json() == before + assert calls + + +def test_plain_eos_metadata_preserves_native_callback_free_path(monkeypatch): + from art.trajectories import _tokenize as module + + class Tokenizer: + eos_token_id = 3 + + def unexpected(*args, **kwargs): + pytest.fail("Plain EOS metadata must not require a callback guard or renderer") + + monkeypatch.setattr(module, "_sampled_source_validator", unexpected) + monkeypatch.setattr(module, "_load_tokenizer", unexpected) + exchange = _chat_exchange([1], [2, 3]) + value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(chat_completions=[exchange])) + actual = value.tokenize(tokenizer=cast(Any, Tokenizer())) + assert actual.tokens == [1, 2, 3] + assert actual.logprobs[-2:] == [-0.2, -0.3] + assert actual.flags[-1] & tr.TokenFlag.STOP + + +def test_recorded_numeric_stop_does_not_read_unused_tokenizer_metadata(): + class Tokenizer: + @property + def eos_token_id(self): + pytest.fail("An explicit recorded STOP needs no tokenizer lookup") + + exchange = _chat_exchange([1], [2, 3]) + assert exchange.response.choices[0].model_extra is not None + exchange.response.choices[0].model_extra["stop_reason"] = 3 + value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(chat_completions=[exchange])) + actual = value.tokenize(tokenizer=cast(Any, Tokenizer())) + assert actual.tokens == [1, 2, 3] + assert actual.flags[-1] & tr.TokenFlag.STOP From 26ffc662d92d77ba7026076a55b7aab9f8d893cb Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 23:22:19 +0000 Subject: [PATCH 30/67] Classify STOP callbacks through the internal rendering wrapper --- src/art/trajectories/_tokenize.py | 2 ++ .../test_stop_attribute_callbacks.py | 21 ++++++++++++++++--- 2 files changed, 20 insertions(+), 3 deletions(-) diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index 86ff0415b..32774affc 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -4962,6 +4962,8 @@ def validate( def _stop_uses_callback(reason: int | str | None, tokenizer: Tokenizer | None) -> bool: + if type(tokenizer) is _RenderingTokenizer: + tokenizer = tokenizer.tokenizer if tokenizer is None or isinstance(reason, int) and not isinstance(reason, bool): return False if isinstance(reason, str) and reason: diff --git a/tests/unit/trajectories/test_stop_attribute_callbacks.py b/tests/unit/trajectories/test_stop_attribute_callbacks.py index 5691ffe0c..3d2fe97b6 100644 --- a/tests/unit/trajectories/test_stop_attribute_callbacks.py +++ b/tests/unit/trajectories/test_stop_attribute_callbacks.py @@ -10,7 +10,10 @@ "attribute", ["eos_token_id", "eot_token_id", "special_tokens_map"] ) @pytest.mark.parametrize("mutate", [False, True]) -def test_stop_metadata_lookup_preserves_consumed_logprobs(attribute, mutate): +@pytest.mark.parametrize("wrapped", [False, True]) +def test_stop_metadata_lookup_preserves_consumed_logprobs(attribute, mutate, wrapped): + from art.trajectories import _tokenize as module + exchange = _chat_exchange([1], [2, 3]) choice = exchange.response.choices[0] assert choice.logprobs is not None and choice.logprobs.content is not None @@ -24,6 +27,10 @@ def read(_self): return {} if attribute == "special_tokens_map" else 3 tokenizer = type("Tokenizer", (), {attribute: property(read)})() + if wrapped: + tokenizer = module._RenderingTokenizer( + cast(Any, tokenizer), module._RenderContextGuard(lambda: []) + ) value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(chat_completions=[exchange])) before = value.model_dump_json() if mutate: @@ -43,7 +50,8 @@ def read(_self): assert calls -def test_plain_eos_metadata_preserves_native_callback_free_path(monkeypatch): +@pytest.mark.parametrize("wrapped", [False, True]) +def test_plain_eos_metadata_preserves_native_callback_free_path(monkeypatch, wrapped): from art.trajectories import _tokenize as module class Tokenizer: @@ -56,7 +64,14 @@ def unexpected(*args, **kwargs): monkeypatch.setattr(module, "_load_tokenizer", unexpected) exchange = _chat_exchange([1], [2, 3]) value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(chat_completions=[exchange])) - actual = value.tokenize(tokenizer=cast(Any, Tokenizer())) + tokenizer = ( + module._RenderingTokenizer( + cast(Any, Tokenizer()), module._RenderContextGuard(lambda: []) + ) + if wrapped + else Tokenizer() + ) + actual = value.tokenize(tokenizer=cast(Any, tokenizer)) assert actual.tokens == [1, 2, 3] assert actual.logprobs[-2:] == [-0.2, -0.3] assert actual.flags[-1] & tr.TokenFlag.STOP From 8a7b725a4b29d0bbcc9ba5a53a79d48381a4ca52 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 23:47:22 +0000 Subject: [PATCH 31/67] Preserve reusable projections and first sufficient native role proofs --- docs/features/additional-histories.mdx | 12 +- src/art/trajectories/_tokenize.py | 202 ++++++++++++------ .../test_historical_prompt_selection.py | 131 ++++++++++++ .../test_plain_eos_opaque_context.py | 91 ++++++++ .../test_reused_responses_projection.py | 41 ++++ 5 files changed, 403 insertions(+), 74 deletions(-) create mode 100644 tests/unit/trajectories/test_historical_prompt_selection.py create mode 100644 tests/unit/trajectories/test_plain_eos_opaque_context.py create mode 100644 tests/unit/trajectories/test_reused_responses_projection.py diff --git a/docs/features/additional-histories.mdx b/docs/features/additional-histories.mdx index 6bbcf71cc..aab3c4501 100644 --- a/docs/features/additional-histories.mdx +++ b/docs/features/additional-histories.mdx @@ -188,8 +188,9 @@ 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. -The role proof covers the complete recorded request containing the last -request-only assistant, including its trailing query. Later sampled native +The role proof uses the first sufficient complete recorded request containing +the last request-only assistant, including its trailing query; this applies to +both normalized and original-template proofs. Later sampled native bodies remain authoritative even when their rendered projections differ. With a supplied renderer, request-owned assistant roles are proved throughout the native stream, including between sampled responses. This requirement applies @@ -222,12 +223,15 @@ When recorded prompt IDs are unavailable, a supplied renderer defines the prompt tokens. ART cannot certify that arbitrary returned IDs implement a particular text transformation. A renderer changing a disposable, single-use message copy does not itself invalidate that authority; changes to borrowed sources or -semantic inputs that ART reuses still require refusal. +semantic inputs that ART reuses still require refusal. In particular, a +Responses message projection reused for prompt and completion rendering must +remain unchanged across those calls. These checks refuse stale results; they do not make callbacks or source objects immutable. Context snapshots reuse shared containers only within one observation; callback-free native assembly retains its bounded response-evidence reuse. Unused opaque request metadata remains valid on complete callback-free native -paths; callbacks require context that ART can compare. +paths, including a supplied tokenizer exposing only plain EOS metadata; +callbacks require context that ART can compare. 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 32774affc..e8c89c2a1 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -1242,10 +1242,13 @@ def set( tokenizer.guard.check() if self.track_sources: self.consume_sources( - sources, require_supported_context=tokenizer is not None + sources, + require_supported_context=_tokenizer_requires_context(tokenizer), ) elif self.validate_sources is not None: - self.validate_sources(None, require_supported_context=tokenizer is not None) + self.validate_sources( + None, require_supported_context=_tokenizer_requires_context(tokenizer) + ) self.tokenizer = ( tokenizer.tokenizer if type(tokenizer) is _RenderingTokenizer else tokenizer ) @@ -3004,12 +3007,24 @@ def _tokenize_exchange_trajectory( ] ledger = _trace or _TraceBuilder(track_sources=False) ledger.consume_sources( - consumed_sources, require_supported_context=tokenizer_instance is not None + consumed_sources, + require_supported_context=_tokenizer_requires_context(tokenizer_instance), ) callback_used = False def checked(function: Callable[..., Any], *args: Any, **kwargs: Any) -> Any: nonlocal callback_used + if ( + function is _template_ids + and (projection := kwargs.get("messages_override")) is not None + ): + # Responses keeps this projection for completion/continuation renders. + # Bind it before the first renderer can change a later ART input. + ledger.consume_auxiliary( + ("response_projection", id(projection)), + lambda projection=projection: projection, + projection, + ) if not callback_used: ledger.consume_sources([], rendered_evidence=True) callback_used = True @@ -3267,7 +3282,9 @@ def fallback_config() -> _TokenizerConfig: ) assert ledger.validate_sources is not None ledger.validate_sources( - None, require_supported_context=callback_used or tokenizer is not None + None, + require_supported_context=callback_used + or _tokenizer_requires_context(tokenizer), ) tokenized = TokenizedHistory( history=history, @@ -4377,7 +4394,7 @@ def read_text( source=source, source_key=source_key, tokenizer=tokenizer, - _callbacks=tokenizer is not None, + _callbacks=_tokenizer_requires_context(tokenizer), ) for offset in range( max(len(token_ids) - len(output), len(token_ids) - stop_count), @@ -4399,12 +4416,16 @@ def read_text( source_keys, sources, tokenizer=tokenizer, - _callbacks=tokenizer is not None, + _callbacks=_tokenizer_requires_context(tokenizer), ) if pending and (ledger.track_sources or ledger.validate_sources is not None): - ledger.consume_sources(pending, require_supported_context=tokenizer is not None) + ledger.consume_sources( + pending, require_supported_context=_tokenizer_requires_context(tokenizer) + ) if ledger.validate_sources is not None: - ledger.validate_sources(None, require_supported_context=tokenizer is not None) + ledger.validate_sources( + None, require_supported_context=_tokenizer_requires_context(tokenizer) + ) tokenized = TokenizedHistory( history=history, model=history.model, @@ -4986,6 +5007,21 @@ def _stop_uses_callback(reason: int | str | None, tokenizer: Tokenizer | None) - return False +def _tokenizer_requires_context(tokenizer: Tokenizer | None) -> bool: + """Plain STOP metadata alone does not introduce an external callback.""" + if type(tokenizer) is _RenderingTokenizer: + tokenizer = tokenizer.tokenizer + if tokenizer is None: + return False + if callable(tokenizer) or _stop_uses_callback(None, tokenizer): + return True + for name in ("apply_chat_template", "get_chat_template", "decode"): + pure, value = _plain_tokenizer_attribute(tokenizer, name) + if not pure or value is not None: + return True + return False + + def _mark_sampled_stops( token_ids: Sequence[int], flags: list[TokenFlag], @@ -7263,63 +7299,81 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: and chat_template is None and chat_template_kwargs is None ): - # Later request-only assistants can have historical rendering - # different from today's normalized template too. Prove the - # entire final recorded request, not only the first prompt. - source = history.message_sources[sampled_message_indices[-1]] - assert source is not None - exchange = _source_exchange(source) - assert isinstance( - exchange, - (ChatCompletionsExchange, MessagesExchange, ResponsesExchange), - ) - prompt = source_prompt_tokens(source) - signature = _source_signature(source) - validate_context = _tokenization_context_validator(history) - masks = None - if prompt and exact.tokens[: len(prompt)] == prompt: - validate_consumed(None) - request_messages, request_tools = _request_messages(exchange) - try: - masks = _recorded_prompt_role_masks( - request_messages, - [None] * len(request_messages), - prompt, - tokenizer=resolved_tokenizer, - template=original_template, - tools=request_tools, - kwargs=kwargs, + # Use the first complete original request covering the final + # request-owned assistant; later native bodies may render differently. + rendering_guard.check() + last_request_assistant = max( + ( + index + for index, (message, source) in enumerate( + zip(messages, history.message_sources, strict=True) ) - except (TypeError, KeyError, NotImplementedError): - pass - validate_context(True) - prompt_cache.clear() - output_cache.clear() - if _source_signature(source) != signature: - raise ValueError( - "Sampled source changed while proving recorded request roles" - ) - if ( - masks is not None - and all( - flag & (TokenFlag.SAMPLED | TokenFlag.OUTPUT) - or not (flag & TokenFlag.ASSISTANT) - or assistant - for flag, assistant in zip(exact.flags, masks[0]) - ) - and all( - flag & (TokenFlag.SAMPLED | TokenFlag.OUTPUT) - or not (flag & TokenFlag.STOP) - or stop - for flag, stop in zip(exact.flags, masks[1]) + if message.get("role") == "assistant" + and (source is None or not _source_is_sampled(source)) + ), + default=-1, + ) + for message_index in sampled_message_indices: + if message_index <= last_request_assistant: + continue + source = history.message_sources[message_index] + assert source is not None + exchange = _source_exchange(source) + assert isinstance( + exchange, + (ChatCompletionsExchange, MessagesExchange, ResponsesExchange), ) - ): - for index, (assistant, stop) in enumerate(zip(*masks, strict=True)): - if not exact.flags[index] & ( - TokenFlag.SAMPLED | TokenFlag.OUTPUT + prompt = source_prompt_tokens(source) + signature = _source_signature(source) + validate_context = _tokenization_context_validator(history) + masks = None + if prompt and exact.tokens[: len(prompt)] == prompt: + validate_consumed(None) + request_messages, request_tools = _request_messages(exchange) + try: + masks = _recorded_prompt_role_masks( + request_messages, + [None] * len(request_messages), + prompt, + tokenizer=resolved_tokenizer, + template=original_template, + tools=request_tools, + kwargs=kwargs, + ) + except (TypeError, KeyError, NotImplementedError): + pass + validate_context(True) + prompt_cache.clear() + output_cache.clear() + if _source_signature(source) != signature: + raise ValueError( + "Sampled source changed while proving recorded request roles" + ) + if ( + masks is not None + and all( + flag & (TokenFlag.SAMPLED | TokenFlag.OUTPUT) + or not (flag & TokenFlag.ASSISTANT) + or assistant + for flag, assistant in zip(exact.flags, masks[0]) + ) + and all( + flag & (TokenFlag.SAMPLED | TokenFlag.OUTPUT) + or not (flag & TokenFlag.STOP) + or stop + for flag, stop in zip(exact.flags, masks[1]) + ) + ): + for index, (assistant, stop) in enumerate( + zip(*masks, strict=True) ): - exact.flags[index] |= _rendered_flag(assistant, False, stop) - return exact + if not exact.flags[index] & ( + TokenFlag.SAMPLED | TokenFlag.OUTPUT + ): + exact.flags[index] |= _rendered_flag( + assistant, False, stop + ) + return exact raise ValueError( "Cannot preserve request roles across exact native prompt replacement" ) @@ -8095,7 +8149,7 @@ def _tokenize_completions_token_history( logprobs[span.start : span.end] = selected_logprobs if consumed: ledger.consume_sources( - consumed, require_supported_context=tokenizer is not None + consumed, require_supported_context=_tokenizer_requires_context(tokenizer) ) _mark_sampled_stops( history.prompt, @@ -8105,7 +8159,9 @@ def _tokenize_completions_token_history( tokenizer=tokenizer, ) if ledger.validate_sources is not None: - ledger.validate_sources(None, require_supported_context=tokenizer is not None) + ledger.validate_sources( + None, require_supported_context=_tokenizer_requires_context(tokenizer) + ) tokenized = TokenizedHistory( history=history, model=history.model, @@ -8321,12 +8377,16 @@ def resolved_tokenizer() -> Tokenizer: source_keys, sources, tokenizer=tokenizer, - _callbacks=tokenizer is not None, + _callbacks=_tokenizer_requires_context(tokenizer), ) if pending and (ledger.track_sources or ledger.validate_sources is not None): - ledger.consume_sources(pending, require_supported_context=tokenizer is not None) + ledger.consume_sources( + pending, require_supported_context=_tokenizer_requires_context(tokenizer) + ) if ledger.validate_sources is not None: - ledger.validate_sources(None, require_supported_context=tokenizer is not None) + ledger.validate_sources( + None, require_supported_context=_tokenizer_requires_context(tokenizer) + ) tokenized = TokenizedHistory( history=history, model=history.model, @@ -8755,7 +8815,8 @@ 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 + builder is not None and _tokenizer_requires_context(builder.tokenizer) + for builder in builders ): # Later callbacks may edit an earlier completed history. Check its # original source keys and stop evidence without calling a tokenizer. @@ -8790,9 +8851,10 @@ def _complete_resolved_sampled_stops( and (tokenizer := resolved.get(value.model)) is not None ): assert builder.validate_sources is not None + requires_context = _tokenizer_requires_context(tokenizer) if builder.validate_context is not None: - builder.validate_context(True) - builder.validate_sources(None) + builder.validate_context(requires_context) + builder.validate_sources(None, require_supported_context=requires_context) _mark_sampled_stops( value.tokens, value.flags, diff --git a/tests/unit/trajectories/test_historical_prompt_selection.py b/tests/unit/trajectories/test_historical_prompt_selection.py new file mode 100644 index 000000000..13bc75a02 --- /dev/null +++ b/tests/unit/trajectories/test_historical_prompt_selection.py @@ -0,0 +1,131 @@ +"""Public three-exchange discriminator; no production monkeypatches.""" + +from copy import deepcopy +import math +from typing import Any, cast + +from openai.types.chat import ChatCompletion +import pytest +from test_recorded_prompt_roles import _interior_historical_case +from test_tokenize import _chat_exchange + +import art.trajectories as tr +from art.trajectories import _tokenize as module + + +def case(native_later=True, mismatch=None, third=True): + value, tokenizer, records = _interior_historical_case() + exchanges = value.exchanges.chat_completions + second = exchanges[1] + prompt = list(records[1][0]) + output = list(map(ord, "native1§" if native_later else "answer1§")) + payload = second.response.model_dump(mode="python") + choice = payload["choices"][0] + choice["token_ids"] = output + choice["logprobs"]["content"] = [ + dict(token=f"token_id:{token}", logprob=-token / 10, bytes=[], top_logprobs=[]) + for token in output + ] + second.response = ChatCompletion.model_validate(payload) + records[1] = (prompt, output) + if mismatch == "query": + second.request["messages"][-1]["content"] = "NEXT QUERY" + elif mismatch == "role": + second.request["messages"][-2]["content"] = "xY" + if third: + final_messages = [ + *deepcopy(second.request["messages"]), + {"role": "assistant", "content": "answer1"}, + {"role": "user", "content": "final query"}, + ] + final_prompt = [*prompt, *output, *map(ord, "final query")] + final_output = list(map(ord, "answer2§")) + last = _chat_exchange(final_prompt, final_output, offset=2) + last.request["messages"] = cast(Any, final_messages) + last.request["chat_template"] = tokenizer.chat_template + payload = last.response.model_dump(mode="python") + payload["choices"][0]["message"]["content"] = "answer2" + last.response = ChatCompletion.model_validate(payload) + exchanges.append(last) + records.append((final_prompt, final_output)) + return value, tokenizer, records + + +def role_proof(exchange, tokenizer): + prompt = exchange.response.choices[0].model_extra["prompt_token_ids"] + return module._recorded_prompt_role_masks( + exchange.request["messages"], + [None] * len(exchange.request["messages"]), + prompt, + tokenizer=tokenizer, + template=tokenizer.chat_template, + tools=None, + kwargs={}, + ) + + +def assert_exact(value, tokenizer, records): + before = value.model_dump_json() + result = value.tokenize(tokenizer=tokenizer, multi_history=True) + assert len(result.histories) == 1 + actual = result.histories[0] + assert actual.tokens == records[-1][0] + records[-1][1] + sampled = ( + tr.TokenFlag.EXACT + | tr.TokenFlag.SAMPLED + | tr.TokenFlag.OUTPUT + | tr.TokenFlag.ASSISTANT + ) + expected_sampled = set() + for prompt, output in records: + assert actual.tokens[: len(prompt)] == prompt + span = range(len(prompt), len(prompt) + len(output)) + expected_sampled.update(span) + assert actual.logprobs[len(prompt) : len(prompt) + len(output)] == [ + -token / 10 for token in output + ] + assert all(actual.flags[i] & sampled == sampled for i in span) + assert actual.flags[len(prompt) + len(output) - 1] & tr.TokenFlag.STOP + assert { + i for i, flag in enumerate(actual.flags) if flag & tr.TokenFlag.SAMPLED + } == expected_sampled + # Original public parser produces y§ for xy. This request-owned + # interior assistant precedes the second query/output; it is never sampled. + start = len(records[0][0]) + len(records[0][1]) + len("middle query") + assert actual.tokens[start : start + 2] == list(map(ord, "y§")) + assert all( + actual.flags[i] & tr.TokenFlag.ASSISTANT for i in range(start, start + 2) + ) + assert actual.flags[start + 1] & tr.TokenFlag.STOP + assert all( + not actual.flags[i] & (tr.TokenFlag.SAMPLED | tr.TokenFlag.OUTPUT) + for i in range(start, start + 2) + ) + assert all(math.isnan(actual.logprobs[i]) for i in range(start, start + 2)) + assert value.model_dump_json() == before + + +def test_first_sufficient_original_request_preserves_later_native_body(): + value, tokenizer, records = case() + # Establish the actual gap before invoking the public route: exchange 2 + # supplies the complete historical role/query proof; exchange 3 cannot + # re-encode the independently exact sampled body native1§ from answer1. + assert role_proof(value.exchanges.chat_completions[1], tokenizer) is not None + assert role_proof(value.exchanges.chat_completions[2], tokenizer) is None + assert_exact(value, tokenizer, records) + + +@pytest.mark.parametrize("third,native_later", [(False, True), (True, False)]) +def test_adjacent_valid_earlier_and_render_congruent_controls(third, native_later): + value, tokenizer, records = case(native_later=native_later, third=third) + assert_exact(value, tokenizer, records) + + +@pytest.mark.parametrize("mismatch", ["query", "role"]) +def test_complete_request_mismatch_still_refuses(mismatch): + value, tokenizer, _ = case(mismatch=mismatch) + before = value.model_dump_json() + assert role_proof(value.exchanges.chat_completions[1], tokenizer) is None + with pytest.raises(ValueError, match="request roles|assistant boundaries"): + value.tokenize(tokenizer=tokenizer, multi_history=True) + assert value.model_dump_json() == before diff --git a/tests/unit/trajectories/test_plain_eos_opaque_context.py b/tests/unit/trajectories/test_plain_eos_opaque_context.py new file mode 100644 index 000000000..d5948a3d1 --- /dev/null +++ b/tests/unit/trajectories/test_plain_eos_opaque_context.py @@ -0,0 +1,91 @@ +from copy import deepcopy +from typing import Any, cast + +import pytest +from test_tokenize import ( + _chat_exchange, + _completion_exchange, + _message_exchange, + _response_exchange, +) + +import art.trajectories as tr + + +@pytest.mark.parametrize("authority", ["none", "plain", "dynamic"]) +@pytest.mark.parametrize( + "protocol", + ["chat_completions", "messages", "responses", "completions", "token_prompt"], +) +def test_callback_free_native_opaque_context_accepts_plain_eos( + authority, protocol, monkeypatch +): + class PlainTokenizer: + eos_token_id = 2 + + class DynamicTokenizer: + @property + def eos_token_id(self): + return 2 + + if protocol == "chat_completions": + exchange = _chat_exchange([1], [2]) + expected_logprob = -0.2 + elif protocol == "messages": + exchange = _message_exchange( + cast( + Any, + dict( + model="test/model", + max_tokens=5, + messages=[dict(role="user", content="question")], + ), + ), + prompt_token_ids=[1], + token_ids=[2], + logprobs=[-0.2], + ) + expected_logprob = -0.2 + elif protocol == "responses": + exchange = _response_exchange("public-opaque", 2, prompt_token_ids=[1]) + expected_logprob = -0.1 + else: + exchange = _completion_exchange( + prompt=[1] if protocol == "token_prompt" else "question" + ) + expected_logprob = -0.2 + protocol = "completions" + opaque = object() + exchange.request["metadata"] = cast(Any, {"unused": opaque}) + # Independent model histories exercise final authority propagation too. + second = deepcopy(exchange) + second.request["model"] = "other/model" + second.request["metadata"] = cast(Any, {"unused": opaque}) + value = tr.Trajectory( + exchanges=tr.TrajectoryExchanges(**{protocol: [exchange, second]}) + ) + + def unexpected(*args, **kwargs): + pytest.fail("Complete native records must not load or render") + + monkeypatch.setattr("art.trajectories._tokenize._load_tokenizer", unexpected) + tokenizer = ( + None + if authority == "none" + else PlainTokenizer() + if authority == "plain" + else DynamicTokenizer() + ) + if authority == "dynamic": + with pytest.raises(ValueError, match="[Cc]ontext"): + value.tokenize(tokenizer=cast(Any, tokenizer), multi_history=True) + return + actual = value.tokenize(tokenizer=cast(Any, tokenizer), multi_history=True) + assert len(actual.histories) == 2 + for result in actual.histories: + assert result.tokens == [1, 2] + assert result.logprobs[-1] == expected_logprob + assert result.flags[-1] & tr.TokenFlag.SAMPLED + assert bool(result.flags[-1] & tr.TokenFlag.STOP) == (authority == "plain") + assert cast(Any, exchange.request["metadata"])["unused"] is opaque + assert cast(Any, second.request["metadata"])["unused"] is opaque diff --git a/tests/unit/trajectories/test_reused_responses_projection.py b/tests/unit/trajectories/test_reused_responses_projection.py new file mode 100644 index 000000000..c32909be9 --- /dev/null +++ b/tests/unit/trajectories/test_reused_responses_projection.py @@ -0,0 +1,41 @@ +from typing import Any, cast + +import pytest +from test_tokenize import _response_exchange + +import art.trajectories as tr + + +@pytest.mark.parametrize("mutate", [False, True]) +def test_responses_renderer_cannot_change_a_reused_projection(mutate): + exchange = _response_exchange("public-reused", 2) + assert exchange.response.model_extra is not None + exchange.response.model_extra.pop("token_generations") + original = exchange.model_dump_json() + observed = [] + + class Tokenizer: + def apply_chat_template(self, messages, **kwargs): + completed = any(message["role"] == "assistant" for message in messages) + observed.append((completed, messages[0]["content"])) + if not completed and mutate: + messages[0]["content"] = "changed detached projection" + return [99, 2] if completed else [99] + + value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(responses=[exchange])) + if mutate: + with pytest.raises( + ValueError, match="[Cc]ontext|[Pp]rojection|Consumed source text" + ): + actual = value.tokenize(tokenizer=cast(Any, Tokenizer())) + assert actual.tokens == [99, 2] + assert observed == [ + (False, "turn 0"), + (True, "changed detached projection"), + ] + assert exchange.model_dump_json() == original + else: + actual = value.tokenize(tokenizer=cast(Any, Tokenizer())) + assert actual.tokens == [99, 2] + assert observed == [(False, "turn 0"), (True, "turn 0")] + assert exchange.model_dump_json() == original From eafc6a989b79c694168c4ecaf13282285c66bfda Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 23:53:34 +0000 Subject: [PATCH 32/67] Retain callback authority through completed history validation --- docs/features/additional-histories.mdx | 2 + src/art/trajectories/_tokenize.py | 11 ++-- .../test_callback_authority_lifetime.py | 64 +++++++++++++++++++ 3 files changed, 73 insertions(+), 4 deletions(-) create mode 100644 tests/unit/trajectories/test_callback_authority_lifetime.py diff --git a/docs/features/additional-histories.mdx b/docs/features/additional-histories.mdx index aab3c4501..280bc3da5 100644 --- a/docs/features/additional-histories.mdx +++ b/docs/features/additional-histories.mdx @@ -208,6 +208,8 @@ already-consumed source tokens, logprobs, model or stop evidence. ART checks eac STOP callback and checks all consumed evidence before returning, including logprobs used only by the rendered fallback. Multi-history tokenization also checks completed histories after later renderer callbacks. +That callback authority is retained even if a callback later removes the +tokenizer methods that made it callback-capable. Evidence is bound when first read, including a later Responses generation read to prove a copied suffix, rendered item IDs and logprobs, and text used for prefix repair. STOP admission also retains the source keys read before its encoder. diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index e8c89c2a1..003f96d50 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -1154,6 +1154,7 @@ class _TraceBuilder: validate_context: Callable[[bool], None] | None = None track_sources: bool = True rendered_evidence: bool = False + callback_authority: bool = False auxiliary_evidence: dict[object, tuple[Callable[[], object], object]] = field( default_factory=dict ) @@ -1216,6 +1217,7 @@ def validate_auxiliary(self) -> None: ) def checked(self, function: Callable[..., Any], *args: Any, **kwargs: Any) -> Any: + self.callback_authority = True if self.validate_context is not None: self.validate_context(True) if self.validate_sources is not None: @@ -1238,6 +1240,7 @@ def set( *, tokenizer: Tokenizer | None = None, ) -> None: + self.callback_authority |= _tokenizer_requires_context(tokenizer) if type(tokenizer) is _RenderingTokenizer: tokenizer.guard.check() if self.track_sources: @@ -8511,6 +8514,9 @@ def _tokenize_history( return _legacy_tokenize(history, model=model) if model is None: raise ValueError("History tokenization requires a model") + if _trace is not None: + # Remember authority before a callback can change its own capabilities. + _trace.callback_authority |= _tokenizer_requires_context(tokenizer) if isinstance(history, CompletionsTokenHistory): return _tokenize_completions_token_history( history, @@ -8814,10 +8820,7 @@ def _materialize_trajectory( def _validate_completed_sources(builders: Sequence[_TraceBuilder | None]) -> None: - if any( - builder is not None and _tokenizer_requires_context(builder.tokenizer) - for builder in builders - ): + if any(builder is not None and builder.callback_authority 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: diff --git a/tests/unit/trajectories/test_callback_authority_lifetime.py b/tests/unit/trajectories/test_callback_authority_lifetime.py new file mode 100644 index 000000000..32ca1d52f --- /dev/null +++ b/tests/unit/trajectories/test_callback_authority_lifetime.py @@ -0,0 +1,64 @@ +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 + + +@pytest.mark.parametrize("mutate", [False, True]) +@pytest.mark.parametrize("internal", [False, True]) +def test_completed_evidence_remains_guarded_after_tokenizer_loses_callbacks( + mutate, internal +): + first = _chat_exchange([1], [2, 9], model="public/first") + second = _chat_exchange([3], [4, 9], model="public/second", offset=1) + assert first.response.choices[0].model_extra is not None + assert second.response.choices[0].model_extra is not None + first.response.choices[0].model_extra["stop_reason"] = 9 + second.response.choices[0].model_extra["stop_reason"] = "§" + calls = [] + lp = first.response.choices[0].logprobs + assert lp is not None and lp.content + first_lp = lp.content + + class PlainEOS: + eos_token_id = 9 + + class Tokenizer(PlainEOS): + def __call__(self, text, **kwargs): + assert text == "§" + calls.append(text) + if mutate: + first_lp[0].logprob = -99 + self.__class__ = PlainEOS + return [9] + + tokenizer = Tokenizer() + value = tr.Trajectory( + exchanges=tr.TrajectoryExchanges(chat_completions=[first, second]) + ) + + def tokenize(): + if internal: + return module._tokenize_trajectory_with_trace( + value, tokenizer=cast(Any, tokenizer) + )[0] + return value.tokenize(tokenizer=cast(Any, tokenizer), multi_history=True) + + if mutate: + with pytest.raises(ValueError, match="Sampled source changed"): + actual = tokenize() + assert actual.histories[0].logprobs[-2] == -0.2 + assert first_lp[0].logprob == -99 + else: + actual = tokenize() + assert [h.tokens for h in actual.histories] == [[1, 2, 9], [3, 4, 9]] + assert [h.logprobs[-2:] for h in actual.histories] == [ + [-0.2, -0.9], + [-0.4, -0.9], + ] + assert all(h.flags[-1] & tr.TokenFlag.STOP for h in actual.histories) + assert calls == ["§"] + assert type(tokenizer) is PlainEOS From 8dff4693987991522451cef0d55381b9d05bd73a Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 27 Sep 2026 00:08:36 +0000 Subject: [PATCH 33/67] Require decoder only for unresolved native boundaries --- src/art/trajectories/_tokenize.py | 6 +- .../test_native_boundary_capabilities.py | 68 +++++++++++++++++++ 2 files changed, 72 insertions(+), 2 deletions(-) create mode 100644 tests/unit/trajectories/test_native_boundary_capabilities.py diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index 003f96d50..6077d5390 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -5610,8 +5610,7 @@ def _tokenize_recorded_chat_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": + if not messages or messages[-1].get("role") != "assistant": return None entries: list[tuple[int, object, list[int], list[int], list[float]]] = [] sources: dict[_SampledSourceKey, object] = {} @@ -5702,6 +5701,9 @@ def decline() -> None: ): continue try: + decode = checked(lambda: getattr(tokenizer, "decode", None)) + if not callable(decode): + return decline() body = checked( lambda: decode( output, diff --git a/tests/unit/trajectories/test_native_boundary_capabilities.py b/tests/unit/trajectories/test_native_boundary_capabilities.py new file mode 100644 index 000000000..4c3eac32b --- /dev/null +++ b/tests/unit/trajectories/test_native_boundary_capabilities.py @@ -0,0 +1,68 @@ +from typing import Any, cast + +import pytest +from test_tokenize import _chat_exchange + +import art.trajectories as tr + + +@pytest.mark.parametrize("decoder", ["absent", "property"]) +@pytest.mark.parametrize( + "case", ["terminal_stop", "terminal_length", "known_prior", "unproved_prior"] +) +def test_native_output_requires_decode_only_for_unproved_nonterminal_boundary( + case, decoder +): + calls = [] + + class Tokenizer: + eos_token_id = 9 + + def __call__(self, *args, **kwargs): + calls.append("encode") + raise RuntimeError("unneeded encoding") + + def apply_chat_template(self, *args, **kwargs): + calls.append("render") + raise RuntimeError("boundary rendering required") + + class DecoderProperty(Tokenizer): + @property + def decode(self): + calls.append("decode") + raise RuntimeError("boundary decoder required") + + if case in ("known_prior", "unproved_prior"): + first_output = [2, 9] if case == "known_prior" else [2] + first = _chat_exchange([1], first_output) + last = _chat_exchange([1, 2, 9, 3], [4], offset=1) + exchanges = [first, last] + expected = [1, 2, 9, 3, 4] + else: + last = _chat_exchange([1], [2]) + if case == "terminal_length": + last.response.choices[0].finish_reason = "length" + exchanges = [last] + expected = [1, 2] + value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(chat_completions=exchanges)) + original = value.model_dump_json() + tokenizer = Tokenizer() if decoder == "absent" else DecoderProperty() + if case == "unproved_prior": + with pytest.raises(RuntimeError, match="boundary .* required"): + value.tokenize(tokenizer=cast(Any, tokenizer)) + assert calls == (["render"] if decoder == "absent" else ["decode"]) + else: + actual = value.tokenize(tokenizer=cast(Any, tokenizer)) + assert actual.tokens == expected + assert actual.logprobs[-1] == -expected[-1] / 10 + assert actual.flags[-1] == ( + tr.TokenFlag.SAMPLED + | tr.TokenFlag.EXACT + | tr.TokenFlag.OUTPUT + | tr.TokenFlag.ASSISTANT + ) + if case == "known_prior": + assert actual.flags[2] & tr.TokenFlag.STOP + assert actual.logprobs[1:3] == [-0.2, -0.9] + assert calls == [] + assert value.model_dump_json() == original From 55823f7e651164f9e8f24d3a71f245a437d5eba2 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 27 Sep 2026 00:11:35 +0000 Subject: [PATCH 34/67] Include scalar and Enum slot storage in context snapshots --- src/art/trajectories/_tokenize.py | 26 +++++- tests/unit/trajectories/test_scalar_slots.py | 95 ++++++++++++++++++++ 2 files changed, 117 insertions(+), 4 deletions(-) create mode 100644 tests/unit/trajectories/test_scalar_slots.py diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index 6077d5390..399496db2 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -17,7 +17,7 @@ from pickle import Pickler, PicklingError import re import threading -from types import FunctionType +from types import FunctionType, MemberDescriptorType from typing import TYPE_CHECKING, Any, Literal, Protocol, cast import warnings @@ -4779,12 +4779,13 @@ def snapshot(item: object) -> object: if isinstance(item, Enum): if getattr(item, "__objclass__", kind) is not kind: raise TypeError("Unsupported enum tokenization context") - return kind, snapshot( + return kind, instance_state( + item, { key: child for key, child in vars(item).items() if key != "__objclass__" - } + }, ) if isinstance(item, (str, int, float, bytes)): if isinstance(item, str): @@ -4795,7 +4796,7 @@ def snapshot(item: object) -> object: scalar = repr(float.__float__(item)) else: scalar = bytes.__bytes__(item) - return kind, scalar, snapshot(getattr(item, "__dict__", None)) + return kind, scalar, instance_state(item, getattr(item, "__dict__", None)) if kind in (list, tuple): result = kind, tuple(snapshot(child) for child in cast(Sequence, item)) elif isinstance(item, Mapping): @@ -4819,6 +4820,23 @@ def snapshot(item: object) -> object: observed[id(item)] = item, result return result + def instance_state(item: object, dictionary: object) -> object: + # Scalar subclasses and Enum members may keep mutable state in slots. + # Read actual slot storage, including inherited/shadowed slots, without + # invoking an instance's attribute lookup or replacement properties. + slots = [] + kind = type(item) + for owner in type.__getattribute__(kind, "__mro__"): + for name, descriptor in type.__getattribute__(owner, "__dict__").items(): + if type(descriptor) is MemberDescriptorType: + try: + child = descriptor.__get__(item, kind) + except AttributeError: + slots.append((owner, name, False)) + else: + slots.append((owner, name, True, snapshot(child))) + return snapshot(dictionary), tuple(slots) + return snapshot(value) diff --git a/tests/unit/trajectories/test_scalar_slots.py b/tests/unit/trajectories/test_scalar_slots.py new file mode 100644 index 000000000..64af949c9 --- /dev/null +++ b/tests/unit/trajectories/test_scalar_slots.py @@ -0,0 +1,95 @@ +from typing import Any, cast + +import pytest +from test_tokenize import _chat_exchange + +import art.trajectories as tr + + +@pytest.mark.parametrize("mutate", [False, True]) +@pytest.mark.parametrize("kind", ["str", "float", "enum"]) +def test_reused_scalar_slot_state_remains_part_of_context(mutate, kind): + from enum import Enum + + class Text(str): + __slots__ = ("state",) + + class Number(float): + __slots__ = ("state",) + + class Choice(Enum): + __slots__ = ("state",) + FIRST = "first" + + setting = cast( + Any, + Text("stable") + if kind == "str" + else Number(1) + if kind == "float" + else Choice.FIRST, + ) + setting.state = [1] + options = {"value": setting} + exchange = _chat_exchange([1], [2]) + exchange.request["messages"] = [{"role": "user", "content": "q"}] + value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(chat_completions=[exchange])) + before = value.model_dump_json() + calls = [] + + class Tokenizer: + def apply_chat_template(self, messages, **kwargs): + assert kwargs["options"] is options + calls.append("render") + if mutate: + setting.state[0] = 3 + return [setting.state[0], 2] + + def __call__(self, text, **kwargs): + calls.append("encode") + return [2 if text == "answer" else setting.state[0]] + + def invoke(): + return value.tokenize( + tokenizer=cast(Any, Tokenizer()), + chat_template="custom", + chat_template_kwargs={"options": options}, + ) + + if mutate: + with pytest.raises(ValueError, match="[Cc]ontext changed"): + result = invoke() + assert result.tokens == [3, 2] + assert result.logprobs[-1] == -0.2 + assert setting.state == [3] + else: + result = invoke() + assert result.tokens == [1, 2] + assert result.logprobs[-1] == -0.2 + assert setting.state == [1] + assert calls + assert value.model_dump_json() == before + + +def test_scalar_slot_snapshot_distinguishes_missing_inherited_and_shadowed_state(): + from art.trajectories._tokenize import _tokenization_context + + class Parent(str): + __slots__ = ("state",) + + class Child(Parent): + __slots__ = ("state",) + + value = Child("stable") + empty = _tokenization_context(value) + vars(Parent)["state"].__set__(value, [1]) + inherited = _tokenization_context(value) + assert empty != inherited + vars(Child)["state"].__set__(value, [2]) + both = _tokenization_context(value) + assert inherited != both + vars(Parent)["state"].__get__(value)[0] = 3 + assert both != _tokenization_context(value) + vars(Parent)["state"].__delete__(value) + vars(Child)["state"].__delete__(value) + assert empty == _tokenization_context(value) From 02ca697e721dbf98703b2a7f24512f162f9ee10b Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 27 Sep 2026 00:12:09 +0000 Subject: [PATCH 35/67] Clarify rendered role extent and restored-edit authority --- docs/features/additional-histories.mdx | 19 ++++++++++++------- 1 file changed, 12 insertions(+), 7 deletions(-) diff --git a/docs/features/additional-histories.mdx b/docs/features/additional-histories.mdx index 280bc3da5..76dc2597a 100644 --- a/docs/features/additional-histories.mdx +++ b/docs/features/additional-histories.mdx @@ -188,10 +188,13 @@ 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. -The role proof uses the first sufficient complete recorded request containing -the last request-only assistant, including its trailing query; this applies to -both normalized and original-template proofs. Later sampled native -bodies remain authoritative even when their rendered projections differ. +The role proof uses a complete recorded request covering the request-only role +spans being attributed, including its trailing query. The normalized path uses +the first sufficient request; the original-template fallback also locates the +last request-only assistant message. Later sampled native bodies remain +authoritative even when their rendered projections differ. An empty assistant +with no rendered token extent supplies no role or STOP authority for unmatched +native tokens; ART leaves those tokens' roles unknown rather than inventing them. With a supplied renderer, request-owned assistant roles are proved throughout the native stream, including between sampled responses. This requirement applies even when template normalization makes no change. A request whose text @@ -224,9 +227,11 @@ callback exceptions keep their original identity. When recorded prompt IDs are unavailable, a supplied renderer defines the prompt tokens. ART cannot certify that arbitrary returned IDs implement a particular text transformation. A renderer changing a disposable, single-use message copy -does not itself invalidate that authority; changes to borrowed sources or -semantic inputs that ART reuses still require refusal. In particular, a -Responses message projection reused for prompt and completion rendering must +does not itself invalidate that authority. ART refuses inconsistent borrowed +provenance and changed semantic inputs that it reuses. An edit to an original +object that is restored before any later consumption, leaves final provenance +unchanged, and changes no tokenization decision is not itself a stale result. +A Responses message projection reused for prompt and completion rendering must remain unchanged across those calls. These checks refuse stale results; they do not make callbacks or source objects immutable. Context snapshots reuse shared containers only within one observation; From eb9e48351b74314bbcff5a646237e22c6566fb0b Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 27 Sep 2026 00:54:32 +0000 Subject: [PATCH 36/67] Preserve recorded fields and request rendering authority --- src/art/trajectories/_tokenize.py | 35 +++- .../test_historical_kwargs_identity.py | 113 +++++++++++++ .../test_initial_empty_request_role.py | 84 ++++++++++ .../test_recorded_logprob_fields.py | 74 +++++++++ .../test_reused_response_selection.py | 157 ++++++++++++++++++ 5 files changed, 457 insertions(+), 6 deletions(-) create mode 100644 tests/unit/trajectories/test_historical_kwargs_identity.py create mode 100644 tests/unit/trajectories/test_initial_empty_request_role.py create mode 100644 tests/unit/trajectories/test_recorded_logprob_fields.py create mode 100644 tests/unit/trajectories/test_reused_response_selection.py diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index 399496db2..677adf55a 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -25,6 +25,7 @@ from openai.types import Completion from openai.types.chat import ChatCompletion from openai.types.chat.chat_completion import Choice +from openai.types.chat.chat_completion_token_logprob import ChatCompletionTokenLogprob from openai.types.responses import Response from pydantic import BaseModel @@ -609,7 +610,8 @@ def _recorded_prompt_tokens( tools: object, kwargs: Mapping[str, object], ) -> list[int]: - messages, tools, kwargs = deepcopy((messages, tools, dict(kwargs))) + messages, tools = deepcopy((messages, tools)) + kwargs = dict(kwargs) context = _render_context_key([messages, tools, kwargs]) def check() -> None: @@ -660,7 +662,8 @@ def _recorded_prompt_role_masks( return None if not any(message.get("role") == "assistant" for message in messages): return None - messages, tools, kwargs = deepcopy((messages, tools, dict(kwargs))) + messages, tools = deepcopy((messages, tools)) + kwargs = dict(kwargs) original_context = _render_context_key([messages, tools, dict(kwargs)]) def check_context() -> None: @@ -1278,7 +1281,7 @@ def _chat_logprob_fingerprint_evidence(choice: Choice) -> dict[str, object] | No def values(items: Sequence[object] | None) -> list[dict[str, object]]: result: list[dict[str, object]] = [] for item in items or []: - data = _dump(item) + data = _token_logprob_data(item) result.append( { key: data[key] @@ -1527,6 +1530,15 @@ def _dump(value: object) -> dict[str, Any]: return _string_dict(value) or {} +def _token_logprob_data(value: object) -> dict[str, Any]: + if isinstance(value, ChatCompletionTokenLogprob): + # Recorded fields, including provider token_id extras, are evidence. + # Serializing a typed row must not execute a model_dump override while + # numeric assembly and its fingerprint consume that same evidence. + return {**value.__dict__, **(value.model_extra or {})} + return _dump(value) + + def _field(value: object, name: str, default: object = None) -> object: return ( value.get(name, default) @@ -1596,7 +1608,7 @@ def _pairs( logprobs: list[float] = [] complete = True for value in values: - data = _dump(value) + data = _token_logprob_data(value) token_id = _pair_token_id(data, required=require_token_ids, field=field) if token_id is None: complete = False @@ -2658,12 +2670,14 @@ def _visible_logprobs( entries = _chat_logprob_entries(choice) decoder = codecs.getincrementaldecoder("utf-8")() for index, entry in enumerate(entries): - data = _dump(entry) + data = _token_logprob_data(entry) raw_bytes = data.get("bytes") if isinstance(raw_bytes, list): try: next_data = ( - _dump(entries[index + 1]) if index + 1 < len(entries) else {} + _token_logprob_data(entries[index + 1]) + if index + 1 < len(entries) + else {} ) text = decoder.decode( bytes(raw_bytes), @@ -3028,6 +3042,12 @@ def checked(function: Callable[..., Any], *args: Any, **kwargs: Any) -> Any: lambda projection=projection: projection, projection, ) + # Selection and rendering are separate callbacks. The retained + # projection must remain stable before either consumes it. + args = ( + _RenderingTokenizer(args[0], _RenderContextGuard(lambda: projection)), + *args[1:], + ) if not callback_used: ledger.consume_sources([], rendered_evidence=True) callback_used = True @@ -6535,6 +6555,9 @@ def matching_source_prompt(source: object) -> list[int] | None: exact_prefix_length = len(source_prompt) canonical_prefix_length = len(rendered_prompt) needs_request_roles = any( + canonical_assistant_mask[:canonical_prefix_length] + + canonical_stop_mask[:canonical_prefix_length] + ) and any( prior_message.get("role") == "assistant" and ( prior_source is None or not _source_is_sampled(prior_source) diff --git a/tests/unit/trajectories/test_historical_kwargs_identity.py b/tests/unit/trajectories/test_historical_kwargs_identity.py new file mode 100644 index 000000000..49311c455 --- /dev/null +++ b/tests/unit/trajectories/test_historical_kwargs_identity.py @@ -0,0 +1,113 @@ +from copy import deepcopy +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 + + +@pytest.mark.parametrize("require_identity", [False, True]) +def test_historical_role_proof_nested_options_identity(require_identity: bool) -> None: + history, tokenizer = _history(content="New recorded answer") + 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": "Public reasoningPublic answer"}, + {"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, + ) + extra = exchange.response.choices[0].model_extra + assert extra is not None + extra["prompt_token_ids"] = prompt + options = {"state": ["unchanged"]} + kwargs = exchange.request["chat_template_kwargs"] + assert isinstance(kwargs, dict) + kwargs["options"] = options + value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(chat_completions=[exchange])) + value = value.chat_completions_history() + assert value.chat_template_kwargs is not None + options = value.chat_template_kwargs["options"] + original = value.model_dump_json() + ordinary = tokenizer.apply_chat_template + seen = [] + + def render(messages, **kwargs: Any): + same = kwargs.get("options") is options + seen.append((kwargs.get("chat_template") == _TEMPLATE, same)) + if require_identity and not same: + raise ValueError("nested option identity changed") + return ordinary(messages, **kwargs) + + tokenizer.apply_chat_template = cast(Any, render) + result = value.tokenize(tokenizer=tokenizer) + assert result.tokens[: len(prompt)] == prompt + assert ( + result.tokens[len(prompt) : len(prompt) + len(extra["token_ids"])] + == extra["token_ids"] + ) + assert value.model_dump_json() == original + assert seen and all(same for _, same in seen) + + +@pytest.mark.parametrize("masks", [False, True]) +@pytest.mark.parametrize("mutate", [False, True]) +def test_historical_helpers_preserve_kwargs_identity_and_checks( + masks: bool, mutate: bool +) -> None: + messages = [{"role": "assistant", "content": "x"}, {"role": "user", "content": "q"}] + tools = [{"type": "function", "function": {"name": "public_tool"}}] + options = {"state": ["unchanged"]} + calls = [] + + class Tokenizer: + def apply_chat_template(self, selected, *, tokenize=True, **kwargs): + assert selected is not messages + assert kwargs["tools"] is not tools + assert kwargs["options"] is options + calls.append(True) + if mutate: + options["state"][0] = "changed" + text = "".join(message["content"] for message in selected) + return list(map(ord, text)) if tokenize else text + + def __call__(self, text, **kwargs): + return { + "input_ids": list(map(ord, text)), + "offset_mapping": [(i, i + 1) for i in range(len(text))], + } + + def invoke(): + kwargs: dict[str, Any] = dict( + tokenizer=Tokenizer(), + template="custom", + tools=tools, + kwargs={"options": options}, + ) + if masks: + return _tokenize._recorded_prompt_role_masks( + messages, [None, None], [120, 113], **kwargs + ) + return _tokenize._recorded_prompt_tokens(messages, **kwargs) + + if mutate: + with pytest.raises(ValueError, match="Renderer changed context"): + invoke() + else: + result = invoke() + assert result == (([True, False], [False, False]) if masks else [120, 113]) + assert calls diff --git a/tests/unit/trajectories/test_initial_empty_request_role.py b/tests/unit/trajectories/test_initial_empty_request_role.py new file mode 100644 index 000000000..7a238affb --- /dev/null +++ b/tests/unit/trajectories/test_initial_empty_request_role.py @@ -0,0 +1,84 @@ +from typing import Any, cast + +import pytest +from test_tokenize import _chat_exchange + +import art.trajectories as tr + + +class _ContentOnly: + eos_token_id: int | None = None + + def __init__(self, assistant_extent: str = "") -> None: + self.assistant_extent = assistant_extent + if assistant_extent: + self.eos_token_id = ord(assistant_extent) + + def __call__(self, text: str, **kwargs: Any) -> dict[str, Any]: + result: dict[str, Any] = {"input_ids": list(map(ord, text))} + if kwargs.get("return_offsets_mapping"): + result["offset_mapping"] = [(i, i + 1) for i in range(len(text))] + return result + + def decode(self, ids: list[int], **kwargs: Any) -> str: + return "".join(map(chr, ids)) + + def apply_chat_template( + self, messages: list[dict[str, Any]], *, tokenize: bool = True, **kwargs: Any + ) -> str | list[int]: + text = "".join( + (message.get("content") or "") + + (self.assistant_extent if message.get("role") == "assistant" else "") + for message in messages + ) + return list(map(ord, text)) if tokenize else text + + +@pytest.mark.parametrize( + "assistant,prompt,extent,accepted", + [ + ("", [113], "", True), + ("", [1, 113], "", True), + (None, [1, 113], "", True), + ("", [72, 113], "H", True), + ("H", [1, 72, 113], "", False), + ], +) +def test_initial_request_role_requires_observed_extent( + assistant: str | None, prompt: list[int], extent: str, accepted: bool +) -> None: + exchange = _chat_exchange(prompt, [97]) + request = [] if assistant is None else [{"role": "assistant", "content": assistant}] + request.append({"role": "user", "content": "q"}) + exchange.request["messages"] = cast(Any, request) + choice = exchange.response.choices[0] + choice.message.content = "a" + choice.finish_reason = "length" + trajectory = tr.Trajectory( + exchanges=tr.TrajectoryExchanges(chat_completions=[exchange]) + ) + history = trajectory.histories()[0] + assert isinstance(history, tr.ChatCompletionsHistory) + before = trajectory.model_dump_json() + tokenizer = _ContentOnly(extent) + if not accepted: + with pytest.raises(ValueError, match="Cannot preserve assistant boundaries"): + history.tokenize(tokenizer=tokenizer) + else: + result = history.tokenize(tokenizer=tokenizer) + assert result.tokens == prompt + [97] + assert result.logprobs[len(prompt) :] == [-9.7] + assert [i for i, f in enumerate(result.flags) if f & tr.TokenFlag.SAMPLED] == [ + len(prompt) + ] + assert [i for i, f in enumerate(result.flags) if f & tr.TokenFlag.OUTPUT] == [ + len(prompt) + ] + assert not result.flags[-1] & tr.TokenFlag.STOP + if extent: + assert result.flags[0] == ( + tr.TokenFlag.EXACT | tr.TokenFlag.ASSISTANT | tr.TokenFlag.STOP + ) + else: + assert result.flags[:-1] == [tr.TokenFlag.EXACT] * len(prompt) + assert trajectory.model_dump_json() == before diff --git a/tests/unit/trajectories/test_recorded_logprob_fields.py b/tests/unit/trajectories/test_recorded_logprob_fields.py new file mode 100644 index 000000000..7b4565bd4 --- /dev/null +++ b/tests/unit/trajectories/test_recorded_logprob_fields.py @@ -0,0 +1,74 @@ +from typing import Any, cast + +from openai.types.chat.chat_completion_token_logprob import ChatCompletionTokenLogprob +import pytest +from test_tokenize import _chat_exchange + +import art.trajectories as tr + + +@pytest.mark.parametrize("kind", ["standard", "pure", "mutating", "raising"]) +@pytest.mark.parametrize("metadata", ["sentinel", "extra_id", "positional"]) +@pytest.mark.parametrize("refusal", [False, True]) +def test_native_logprobs_read_recorded_fields( + kind: str, metadata: str, refusal: bool +) -> None: + exchange = _chat_exchange([1], [2]) + choice = exchange.response.choices[0] + assert choice.logprobs is not None and choice.logprobs.content + calls: list[float] = [] + armed = False + + class Entry(ChatCompletionTokenLogprob): + def model_dump(self, *args: Any, **kwargs: Any) -> dict[str, Any]: + result = super().model_dump(*args, **kwargs) + if armed: + calls.append(self.logprob) + if kind == "raising": + raise RuntimeError("Numeric evidence must not call serialization") + if kind == "mutating" and len(calls) == 2: + self.logprob = -9.0 + return result + + row = choice.logprobs.content[0] + if kind != "standard": + row = Entry.model_validate(row.model_dump()) + if metadata != "sentinel": + row.token = "visible text" + if metadata == "extra_id": + assert row.model_extra is not None + row.model_extra["token_id"] = 2 + choice.logprobs.content = [] if refusal else [row] + choice.logprobs.refusal = [row] if refusal else None + trajectory = tr.Trajectory( + exchanges=tr.TrajectoryExchanges(chat_completions=[exchange]) + ) + history = trajectory.chat_completions_history() + before = trajectory.model_dump_json() + armed = True + result = history.tokenize() + assert calls == [] + assert result.tokens == [1, 2] + assert result.logprobs[-1] == row.logprob == -0.2 + assert result.flags == [ + tr.TokenFlag.EXACT, + tr.TokenFlag.EXACT + | tr.TokenFlag.ASSISTANT + | tr.TokenFlag.OUTPUT + | tr.TokenFlag.SAMPLED, + ] + assert trajectory.model_dump_json() == before + + +@pytest.mark.parametrize("token_id", [True, -1, "invalid", 3]) +def test_native_typed_extra_token_id_keeps_validation(token_id: object) -> None: + exchange = _chat_exchange([1], [2]) + choice = exchange.response.choices[0] + assert choice.logprobs is not None and choice.logprobs.content + row = choice.logprobs.content[0] + cast(dict[str, Any], row.model_extra)["token_id"] = token_id + trajectory = tr.Trajectory( + exchanges=tr.TrajectoryExchanges(chat_completions=[exchange]) + ) + with pytest.raises(ValueError, match="invalid exact token ID|disagree"): + trajectory.tokenize() diff --git a/tests/unit/trajectories/test_reused_response_selection.py b/tests/unit/trajectories/test_reused_response_selection.py new file mode 100644 index 000000000..1cdb0cdca --- /dev/null +++ b/tests/unit/trajectories/test_reused_response_selection.py @@ -0,0 +1,157 @@ +from __future__ import annotations + +import hashlib +import math +from typing import Any, cast + +from openai.types.responses import Response +import pytest +from test_tokenize import _response_exchange + +import art.trajectories as tr + + +def _run_case(mode: str, behavior: str) -> dict[str, Any]: + exchange = _response_exchange("public-selection-disposition", 2) + data = exchange.response.model_dump(mode="python") + data.pop("token_generations") + data["output"][0]["content"][0]["logprobs"] = [ + { + "token": "answer", + "logprob": -0.2, + "bytes": list(b"answer"), + "top_logprobs": [], + } + ] + exchange.response = Response.model_validate(data) + original = exchange.model_dump_json() + calls = [] + + class Tokenizer: + chat_template = {"default": "public template"} + retained = None + + def get_chat_template(self, **kwargs): + before = self.retained["content"] if self.retained is not None else None + if self.retained is not None and mode != "none": + self.retained["content"] = "changed before completion" + if mode == "restored_before_consumer": + self.retained["content"] = "turn 0" + calls.append( + { + "call": "select", + "before": before, + "after": self.retained["content"] + if self.retained is not None + else None, + } + ) + return "public template" + + def apply_chat_template(self, messages, **kwargs): + completed = any(message["role"] == "assistant" for message in messages) + consumed = messages[0]["content"] + changed = consumed != "turn 0" + prompt = 101 if completed and changed and behavior == "prefix" else 99 + output = 3 if changed and behavior == "output" else 2 + tokens = [prompt, output] if completed else [prompt] + calls.append( + { + "call": "render", + "completed": completed, + "consumed": consumed, + "same_retained_object": self.retained is messages[0], + "returned": tokens, + } + ) + if not completed: + self.retained = messages[0] + elif mode == "consumed_then_restored": + messages[0]["content"] = "turn 0" + return tokens + + def __call__(self, text, **kwargs): + calls.append( + { + "call": "encode", + "text": text, + "offsets": bool(kwargs.get("return_offsets_mapping")), + } + ) + assert text == "answer" + return ( + {"input_ids": [2], "offset_mapping": [(0, len(text))]} + if kwargs.get("return_offsets_mapping") + else [2] + ) + + tokenizer = Tokenizer() + value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(responses=[exchange])) + row = { + "mode": mode, + "behavior": behavior, + "native_prompt_ids": None, + "native_output_ids": None, + "native_token_generations": None, + "visible_response_text": "answer", + "visible_response_logprob": -0.2, + } + try: + result = value.tokenize(tokenizer=cast(Any, tokenizer)) + row["result"] = { + "tokens": list(result.tokens), + "logprobs": [None if math.isnan(x) else x for x in result.logprobs], + "flags": [int(x) for x in result.flags], + "first_masks": { + flag.name: tr.first_occurrence_masks([result], where=flag) + for flag in ( + tr.TokenFlag.SAMPLED, + tr.TokenFlag.OUTPUT, + tr.TokenFlag.ASSISTANT, + tr.TokenFlag.STOP, + ) + }, + } + except Exception as exc: + row["error"] = {"class": type(exc).__name__, "message": str(exc)} + row.update( + calls=calls, + original_exchange_unchanged=exchange.model_dump_json() == original, + original_exchange_sha256=hashlib.sha256(original.encode()).hexdigest(), + retained_final=tokenizer.retained["content"] + if tokenizer.retained is not None + else None, + ) + return row + + +@pytest.mark.parametrize("behavior", ["same", "prefix", "output"]) +@pytest.mark.parametrize( + "mode", ["none", "restored_before_consumer", "consumed_then_restored", "lasting"] +) +def test_responses_selection_does_not_change_reused_projection( + mode: str, behavior: str +) -> None: + row = _run_case(mode, behavior) + assert row["original_exchange_unchanged"] + rendered = [call for call in row["calls"] if call["call"] == "render"] + assert rendered[0]["consumed"] == "turn 0" + if mode in ("consumed_then_restored", "lasting"): + assert row["error"]["class"] == "ValueError" + assert "context changed" in row["error"]["message"] + # The changed projection never reaches the completion renderer. + assert len(rendered) == 1 + else: + assert len(rendered) == 2 + assert rendered[1]["consumed"] == "turn 0" + assert rendered[1]["same_retained_object"] + assert row["result"]["tokens"] == [99, 2] + assert row["result"]["logprobs"] == [None, -0.2] + assert row["result"]["flags"] == [0, 20] + assert row["result"]["first_masks"] == { + "SAMPLED": [[False, False]], + "OUTPUT": [[False, True]], + "ASSISTANT": [[False, True]], + "STOP": [[False, False]], + } + assert row["retained_final"] == "turn 0" From 0fc97edb27ce0f32258553e9238d66e29a8ccc22 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 27 Sep 2026 02:45:56 +0000 Subject: [PATCH 37/67] Reuse unchanged evidence serialization within source validation --- src/art/trajectories/_tokenize.py | 62 +++++- .../test_evidence_serialization_cache.py | 197 ++++++++++++++++++ 2 files changed, 252 insertions(+), 7 deletions(-) create mode 100644 tests/unit/trajectories/test_evidence_serialization_cache.py diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index 677adf55a..b137dcbc3 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -1274,6 +1274,38 @@ def _fingerprint(value: object) -> str: return sha256(serialized.encode()).hexdigest() +class _RevalidatedFingerprints(dict[tuple[int, str, int], tuple[Exchange, str]]): + """Reuse JSON encoding only after freshly comparing detached evidence.""" + + def __init__(self) -> None: + self.observations: dict[tuple[int, str, int], tuple[Exchange, bytes, str]] = {} + self.snapshot_bytes = 0 + + def fingerprint( + self, exchange: Exchange, protocol: str, index: int, evidence: object + ) -> str: + if type(protocol) is not str or type(index) is not int: + return _fingerprint(evidence) + try: + # Exact builtins only: no user reducers, alias-dependent memo or + # borrowed containers. Float bits and primitive types stay distinct. + snapshot = _RenderContextGuard.plain(evidence, memo=False) + except (TypeError, ValueError, RecursionError, PicklingError): + return _fingerprint(evidence) + key = id(exchange), protocol, index + previous = self.observations.get(key) + if previous is not None and previous[0] is exchange and previous[1] == snapshot: + return previous[2] + result = _fingerprint(evidence) + retained_bytes = self.snapshot_bytes - (len(previous[1]) if previous else 0) + if ( + previous is not None or len(self.observations) < 256 + ) and retained_bytes + len(snapshot) <= 8 << 20: + self.observations[key] = exchange, snapshot, result + self.snapshot_bytes = retained_bytes + len(snapshot) + return result + + def _chat_logprob_fingerprint_evidence(choice: Choice) -> dict[str, object] | None: if choice.logprobs is None: return None @@ -1304,7 +1336,7 @@ def _sampled_evidence_fingerprint( index: int, _cache: dict[tuple[int, str, int], tuple[Exchange, str]] | None = None, ) -> str: - if _cache is not None: + if _cache is not None and type(_cache) is not _RevalidatedFingerprints: key = (id(exchange), protocol, index) cached = _cache.get(key) if cached is not None and cached[0] is exchange: @@ -1400,7 +1432,11 @@ def _sampled_evidence_fingerprint( "logprobs": response_extra.get("logprobs"), "stop_reason": exchange.response.stop_reason, } - return _fingerprint(evidence) + return ( + _cache.fingerprint(exchange, protocol, index, evidence) + if type(_cache) is _RevalidatedFingerprints + else _fingerprint(evidence) + ) def _source_key( @@ -1475,12 +1511,17 @@ def _sampled_source_key( raise ValueError("Sampled token source has an unsupported exchange") -def _exchange_sampled_source_key(exchange: Exchange) -> _SampledSourceKey: +def _exchange_sampled_source_key( + exchange: Exchange, + *, + _fingerprints: dict[tuple[int, str, int], tuple[Exchange, str]] | None = None, +) -> _SampledSourceKey: if isinstance(exchange, ChatCompletionsExchange): return _source_key( exchange, protocol="chat_completions", index=exchange.response.choices[0].index, + _fingerprints=_fingerprints, ) if isinstance(exchange, CompletionsExchange): return _source_key( @@ -1488,11 +1529,16 @@ def _exchange_sampled_source_key(exchange: Exchange) -> _SampledSourceKey: protocol="completions", index=exchange.response.choices[0].index, prompt_index=0, + _fingerprints=_fingerprints, ) if isinstance(exchange, ResponsesExchange): - return _source_key(exchange, protocol="responses", index=0) + return _source_key( + exchange, protocol="responses", index=0, _fingerprints=_fingerprints + ) if isinstance(exchange, MessagesExchange): - return _source_key(exchange, protocol="messages", index=0) + return _source_key( + exchange, protocol="messages", index=0, _fingerprints=_fingerprints + ) raise TypeError(f"Unsupported sampled exchange: {type(exchange).__name__}") @@ -4947,6 +4993,8 @@ def _sampled_source_validator( ) ) + fingerprints = _RevalidatedFingerprints() + def validate( selected: _SampledSourceKey | None, *, @@ -4967,9 +5015,9 @@ def validate( visible, ) in expected[key]: current = ( - _exchange_sampled_source_key(source) + _exchange_sampled_source_key(source, _fingerprints=fingerprints) if isinstance(source, Exchange) - else _sampled_source_key(source) + else _sampled_source_key(source, _fingerprints=fingerprints) ) if ( current != key diff --git a/tests/unit/trajectories/test_evidence_serialization_cache.py b/tests/unit/trajectories/test_evidence_serialization_cache.py new file mode 100644 index 000000000..e64334568 --- /dev/null +++ b/tests/unit/trajectories/test_evidence_serialization_cache.py @@ -0,0 +1,197 @@ +from __future__ import annotations + +import copy +from typing import Any, cast + +from openai.types.chat.chat_completion_message import ChatCompletionMessage +import pytest +from test_tokenize import _chat_exchange + +import art.trajectories._tokenize as core + + +@pytest.mark.parametrize( + ("before", "after"), [(0.0, -0.0), (1, 1.0), (True, 1), (-0.2, -9.0)] +) +@pytest.mark.parametrize("selected", [False, True]) +def test_each_guard_still_detects_changed_primitive_evidence(before, after, selected): + source = cast(Any, _chat_exchange([1], [2])) + row = source.response.choices[0].logprobs.content[0] + row.logprob = before + key = core._exchange_sampled_source_key(source) + validate = core._sampled_source_validator({key: source}) + validate(None) + validate(None) + row.logprob = after + assert core._exchange_sampled_source_key(source) != key + with pytest.raises(ValueError, match="Sampled source changed"): + validate(key if selected else None) + + +@pytest.mark.parametrize("which", [0, 1]) +@pytest.mark.parametrize("field", ["prompt", "output", "bytes", "message"]) +def test_equal_key_aliases_recheck_mutable_storage(which, field): + first = cast(Any, _chat_exchange([1], [2])) + sources = [first, copy.deepcopy(first)] + key = core._exchange_sampled_source_key(first) + validate = core._sampled_source_validator([(key, source) for source in sources]) + validate(None) + validate(None) + source = sources[which] + choice = source.response.choices[0] + if field == "prompt": + choice.model_extra["prompt_token_ids"][0] = 9 + elif field == "output": + choice.model_extra["token_ids"][0] = 9 + elif field == "bytes": + choice.logprobs.content[0].bytes = [9] + else: + choice.message.content = "edited" + with pytest.raises(ValueError, match="Sampled source changed"): + validate(None) + + +def test_message_serialization_callback_still_runs_at_every_guard(): + calls = [] + source = cast(Any, _chat_exchange([1], [2])) + + class Message(ChatCompletionMessage): + def model_dump(self, *args: Any, **kwargs: Any) -> dict[str, Any]: + calls.append(True) + return super().model_dump(*args, **kwargs) + + choice = source.response.choices[0] + choice.message = Message.model_validate(choice.message.model_dump()) + key = core._exchange_sampled_source_key(source) + validate = core._sampled_source_validator({key: source}) + calls.clear() + for _ in range(3): + validate(None) + assert len(calls) == 3 + + +def test_fresh_equal_values_reuse_json_only_not_evidence_reads(monkeypatch): + source = cast(Any, _chat_exchange([1], [2])) + key = core._exchange_sampled_source_key(source) + validate = core._sampled_source_validator({key: source}) + original = core._fingerprint + calls = [] + + def record(value): + calls.append(True) + return original(value) + + monkeypatch.setattr(core, "_fingerprint", record) + validate(None) + source.response = copy.deepcopy(source.response) + validate(None) + validate(None) + assert len(calls) == 1 + source.response.choices[0].logprobs.content[0].logprob = -9.0 + with pytest.raises(ValueError, match="Sampled source changed"): + validate(None) + assert len(calls) == 2 + + +def test_rich_values_preserve_original_json_without_executing_reducers(monkeypatch): + reductions = [] + + class Number(float): + def __reduce_ex__(self, protocol): + reductions.append(protocol) + raise AssertionError("must not run reducer") + + source = cast(Any, _chat_exchange([1], [2])) + row = source.response.choices[0].logprobs.content[0] + row.logprob = Number(-0.2) + key = core._exchange_sampled_source_key(source) + validate = core._sampled_source_validator({key: source}) + original = core._fingerprint + calls = [] + + def record(value): + calls.append(True) + return original(value) + + monkeypatch.setattr(core, "_fingerprint", record) + validate(None) + validate(None) + assert len(calls) == 2 and reductions == [] + row.logprob = Number(-9.0) + with pytest.raises(ValueError, match="Sampled source changed"): + validate(None) + assert reductions == [] + + +@pytest.mark.parametrize("evidence", [{"bad": b"bytes"}, {"bad": object()}]) +def test_original_json_errors_are_not_hidden(evidence): + cache = core._RevalidatedFingerprints() + source = cast(Any, _chat_exchange([1], [2])) + with pytest.raises(TypeError) as expected: + core._fingerprint(evidence) + with pytest.raises(TypeError) as actual: + cache.fingerprint(source, "chat_completions", 0, evidence) + assert str(actual.value) == str(expected.value) + assert not cache.observations + + +def test_cycles_fall_back_to_original_error(): + evidence: dict[str, Any] = {} + evidence["cycle"] = evidence + with pytest.raises(ValueError, match="Circular reference detected"): + core._RevalidatedFingerprints().fingerprint( + _chat_exchange([1], [2]), "chat_completions", 0, evidence + ) + + +def test_observation_bound_and_value_semantics(): + cache = core._RevalidatedFingerprints() + sources = [_chat_exchange([1], [2], offset=i) for i in range(257)] + for source in sources: + assert cache.fingerprint( + source, "chat_completions", 0, {"a": [1]} + ) == core._fingerprint({"a": [1]}) + assert len(cache.observations) == 256 + # Distinct layouts/order are permitted by JSON value equality; a miss must + # recompute, never reject an otherwise unchanged source. + shared = [1] + evidence = {"a": shared, "b": shared} + expected = cache.fingerprint(sources[0], "chat_completions", 0, evidence) + assert ( + cache.fingerprint(sources[0], "chat_completions", 0, {"b": [1], "a": [1]}) + == expected + ) + assert ( + cache.fingerprint(sources[0], "chat_completions", 0, {"b": [1], "a": [2]}) + != expected + ) + + +def test_large_valid_evidence_keeps_original_result_without_retention(): + cache = core._RevalidatedFingerprints() + source = _chat_exchange([1], [2]) + evidence = {"text": "x" * (9 << 20)} + expected = core._fingerprint(evidence) + assert cache.fingerprint(source, "chat_completions", 0, evidence) == expected + assert not cache.observations and cache.snapshot_bytes == 0 + + +def test_rich_key_metadata_does_not_add_hash_or_equality_callbacks(): + calls = [] + + class Index(int): + def __hash__(self): + calls.append("hash") + raise AssertionError("cache must not hash rich key") + + def __eq__(self, other): + calls.append("eq") + raise AssertionError("cache must not compare rich key") + + cache = core._RevalidatedFingerprints() + source = _chat_exchange([1], [2]) + for _ in range(2): + assert cache.fingerprint( + source, "chat_completions", Index(0), {"x": 1} + ) == core._fingerprint({"x": 1}) + assert calls == [] and not cache.observations From 5f818a3e94b6f48aae89810f1ecc62db37f22eb4 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 06:44:22 +0000 Subject: [PATCH 38/67] Reuse revalidated fingerprints across tokenization phases --- src/art/trajectories/_tokenize.py | 31 ++++--- .../test_builder_fingerprint_cache.py | 83 +++++++++++++++++++ 2 files changed, 104 insertions(+), 10 deletions(-) create mode 100644 tests/unit/trajectories/test_builder_fingerprint_cache.py diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index b137dcbc3..b97ec561f 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -1164,6 +1164,9 @@ class _TraceBuilder: consumed_sources: dict[ tuple[_SampledSourceKey, int], tuple[_SampledSourceKey, object] ] = field(default_factory=dict) + fingerprints: _RevalidatedFingerprints = field( + default_factory=lambda: _RevalidatedFingerprints() + ) def consume_sources( self, @@ -1202,6 +1205,7 @@ def consume_sources( list(self.consumed_sources.values()), selected_request_fields=selected_request_fields, rendered_evidence=self.rendered_evidence, + _fingerprints=self.fingerprints, ) def consume_auxiliary( @@ -3921,7 +3925,7 @@ def flush() -> None: 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=ledger.fingerprints) if source_key in seen: continue seen.add(source_key) @@ -4947,6 +4951,7 @@ def _sampled_source_validator( *, selected_request_fields: tuple[str, ...] | None = None, rendered_evidence: bool = False, + _fingerprints: _RevalidatedFingerprints | None = None, ) -> _SampledSourceValidator: expected = {} observed: dict[int, tuple[object, object]] = {} @@ -4993,7 +4998,9 @@ def _sampled_source_validator( ) ) - fingerprints = _RevalidatedFingerprints() + fingerprints = ( + _RevalidatedFingerprints() if _fingerprints is None else _fingerprints + ) def validate( selected: _SampledSourceKey | None, @@ -5700,6 +5707,9 @@ def _tokenize_recorded_chat_boundaries( return None entries: list[tuple[int, object, list[int], list[int], list[float]]] = [] sources: dict[_SampledSourceKey, object] = {} + fingerprints = ( + _trace.fingerprints if _trace is not None else _RevalidatedFingerprints() + ) for index, (message, source) in enumerate( zip(messages, history.message_sources, strict=True) ): @@ -5707,7 +5717,7 @@ def _tokenize_recorded_chat_boundaries( continue if source is None or not _source_is_sampled(source): return None - key = _sampled_source_key(source) + key = _sampled_source_key(source, _fingerprints=fingerprints) prompt, output, logprobs = _chat_source_record(source) if key in sources or prompt is None or output is None: return None @@ -5741,7 +5751,7 @@ def _tokenize_recorded_chat_boundaries( boundaries: dict[_SampledSourceKey, _RenderedLengthStopBoundary] = {} # Entries already consumed every native record. No callback may replace # that evidence before a later source or the final assembler reads it again. - validate_sources = _sampled_source_validator(sources) + validate_sources = _sampled_source_validator(sources, _fingerprints=fingerprints) validate_context = _tokenization_context_validator(history) keys = tuple(sources) selected_key: _SampledSourceKey | None = None @@ -6098,7 +6108,7 @@ def _tokenize_chat_view( ) # Current projected sources have already been inspected for admission. consumed_keys = { - id(source): _sampled_source_key(source) + id(source): _sampled_source_key(source, _fingerprints=ledger.fingerprints) for source in history.message_sources if source is not None and _source_is_sampled(source) } @@ -8645,13 +8655,14 @@ def _tokenize_history( if _projection_validated else _history_render_state(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. + # Only a callback-free phase can memoize by identity. With a tokenizer, + # retain the encoding cache while freshly rereading evidence at every use. fingerprints: dict[tuple[int, str, int], tuple[Exchange, str]] | None = ( {} if tokenizer is None and isinstance(history, ChatCompletionsHistory) - else None + else _trace.fingerprints + if _trace is not None + else _RevalidatedFingerprints() ) has_length_stop = can_render and _history_has_length_stop( history, _fingerprints=fingerprints @@ -8704,7 +8715,7 @@ def _tokenize_history( _trace=_trace, _strict_sources=True, _prior=_prior, - _fingerprints=fingerprints, + _fingerprints=fingerprints if tokenizer is None else None, ) ) ): diff --git a/tests/unit/trajectories/test_builder_fingerprint_cache.py b/tests/unit/trajectories/test_builder_fingerprint_cache.py new file mode 100644 index 000000000..98ec20eca --- /dev/null +++ b/tests/unit/trajectories/test_builder_fingerprint_cache.py @@ -0,0 +1,83 @@ +from __future__ import annotations + +import copy +from typing import Any, cast + +import pytest +from test_tokenize import _chat_exchange + +import art.trajectories._tokenize as core + + +def test_incremental_consumption_reuses_only_validated_serialization(monkeypatch): + sources = [_chat_exchange([1], [2], offset=index) for index in range(3)] + keys = [core._exchange_sampled_source_key(source) for source in sources] + fingerprint = core._fingerprint + encoded = [] + + def observe(value): + encoded.append(True) + return fingerprint(value) + + monkeypatch.setattr(core, "_fingerprint", observe) + ledger = core._TraceBuilder() + for key, source in zip(keys, sources, strict=True): + ledger.consume_sources({key: source}) + ledger.checked(lambda: None) + assert len(encoded) == len(sources) + # Adding renderer authority rebuilds the guard, not its stable encoding. + ledger.consume_sources([], rendered_evidence=True) + ledger.checked(lambda: None) + assert len(encoded) == len(sources) + + +@pytest.mark.parametrize( + "field", ["prompt", "output", "bytes", "logprob", "request", "stop", "model"] +) +def test_extension_does_not_bless_mutation_of_earlier_sources(field): + first = cast(Any, _chat_exchange([1], [2])) + second = _chat_exchange([1], [3], offset=1) + ledger = core._TraceBuilder() + ledger.consume_sources({core._exchange_sampled_source_key(first): first}) + ledger.checked(lambda: None) + ledger.consume_sources({core._exchange_sampled_source_key(second): second}) + ledger.checked(lambda: None) + choice = first.response.choices[0] + if field == "prompt": + choice.model_extra["prompt_token_ids"][0] = 8 + elif field == "output": + choice.model_extra["token_ids"][0] = 8 + elif field == "bytes": + choice.logprobs.content[0].bytes = [8] + elif field == "logprob": + choice.logprobs.content[0].logprob = -9.0 + elif field == "request": + first.request["messages"] = [{"role": "user", "content": "changed"}] + elif field == "stop": + choice.model_extra["stop_reason"] = 999 + else: + first.request["model"] = "changed" + with pytest.raises(ValueError, match="changed during tokenization callback"): + ledger.checked(lambda: None) + + +def test_aliases_keep_independent_observations_after_extension(): + first = cast(Any, _chat_exchange([1], [2])) + second = copy.deepcopy(first) + key = core._exchange_sampled_source_key(first) + ledger = core._TraceBuilder() + ledger.consume_sources([(key, first)]) + ledger.checked(lambda: None) + ledger.consume_sources([(key, second)]) + ledger.checked(lambda: None) + assert len(ledger.fingerprints.observations) == 2 + second.response.choices[0].logprobs.content[0].bytes = [8] + with pytest.raises(ValueError, match="Sampled source changed"): + ledger.checked(lambda: None) + + +def test_distinct_builders_do_not_share_observations(): + first, second = core._TraceBuilder(), core._TraceBuilder() + assert first.fingerprints is not second.fingerprints + assert not first.fingerprints.observations + assert not second.fingerprints.observations From aa397391acc73118c709bea3ef9347c8cdb942d4 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 07:40:55 +0000 Subject: [PATCH 39/67] Avoid duplicate identity work in tokenization context snapshots --- src/art/trajectories/_tokenize.py | 11 ++-- .../trajectories/test_context_observation.py | 56 +++++++++++++++++++ 2 files changed, 62 insertions(+), 5 deletions(-) create mode 100644 tests/unit/trajectories/test_context_observation.py diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index b97ec561f..d90026e8f 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -4839,12 +4839,13 @@ def _tokenization_context( def snapshot(item: object) -> object: kind = type(item) - previous = observed.get(id(item)) + identity = id(item) + previous = observed.get(identity) if previous is not None and previous[0] is item: return previous[1] if kind in (str, int, bool, bytes, type(None), datetime, float): result = kind, repr(item) if kind is float else item - observed[id(item)] = item, result + observed[identity] = item, result return result if isinstance(item, Enum): if getattr(item, "__objclass__", kind) is not kind: @@ -4869,13 +4870,13 @@ def snapshot(item: object) -> object: return kind, scalar, instance_state(item, getattr(item, "__dict__", None)) if kind in (list, tuple): result = kind, tuple(snapshot(child) for child in cast(Sequence, item)) - elif isinstance(item, Mapping): + elif kind is dict or isinstance(item, Mapping): result = ( kind, tuple((snapshot(key), snapshot(child)) for key, child in item.items()), ) elif isinstance(item, Exchange): - result = kind, id(item), item.model, snapshot(item.request) + result = kind, identity, item.model, snapshot(item.request) elif isinstance(item, BaseModel): result = ( kind, @@ -4887,7 +4888,7 @@ def snapshot(item: object) -> object: ) else: raise TypeError("Unsupported mutable tokenization context") - observed[id(item)] = item, result + observed[identity] = item, result return result def instance_state(item: object, dictionary: object) -> object: diff --git a/tests/unit/trajectories/test_context_observation.py b/tests/unit/trajectories/test_context_observation.py new file mode 100644 index 000000000..e70d94ece --- /dev/null +++ b/tests/unit/trajectories/test_context_observation.py @@ -0,0 +1,56 @@ +from collections import UserDict + +import pytest + +from art.trajectories._tokenize import _tokenization_context + + +class ObservedDict(dict): + def items(self): + self.reads += 1 + return super().items() + + +@pytest.mark.parametrize("mapping_type", [dict, UserDict, ObservedDict]) +def test_context_preserves_mapping_order_aliases_and_fresh_mutable_reads(mapping_type): + child = ["before"] + mapping = mapping_type({"first": child, "second": child}) + if isinstance(mapping, ObservedDict): + mapping.reads = 0 + original = _tokenization_context([mapping, mapping]) + first, second = original[1] + assert first is second + assert first[0] is mapping_type + assert [entry[0] for entry in first[1]] == [(str, "first"), (str, "second")] + assert first[1][0][1] is first[1][1][1] + if isinstance(mapping, ObservedDict): + assert mapping.reads == 1 + child[0] = "after" + assert _tokenization_context([mapping, mapping]) != original + child[0] = "before" + assert _tokenization_context([mapping, mapping]) == original + assert mapping["first"] is mapping["second"] is child + + +def test_context_rejects_stale_identity_entry_without_borrowing_its_value(): + value = {"request": [1, True, "1"]} + observed = {id(value): (object(), "stale")} + assert _tokenization_context(value, _observed=observed) == _tokenization_context(value) + assert observed[id(value)][0] is value + + +@pytest.mark.parametrize("left,right", [(1, True), (0.0, -0.0), ("x", b"x"), (None, "None")]) +def test_context_retains_typed_scalar_distinctions(left, right): + assert _tokenization_context({"value": left}) != _tokenization_context({"value": right}) + + +def test_context_preserves_custom_mapping_exception_identity(): + failure = RuntimeError("mapping observation failed") + + class Broken(dict): + def items(self): + raise failure + + with pytest.raises(RuntimeError) as caught: + _tokenization_context(Broken(value=1)) + assert caught.value is failure From ea9e74769c98aa5719433c5e75766ba1c0efd4c9 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 08:26:41 +0000 Subject: [PATCH 40/67] Inline exact logprob snapshots while preserving observation order --- src/art/trajectories/_tokenize.py | 6 +- .../trajectories/test_logprob_projection.py | 217 ++++++++++++++++++ 2 files changed, 222 insertions(+), 1 deletion(-) create mode 100644 tests/unit/trajectories/test_logprob_projection.py diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index d90026e8f..2690ceca2 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -1317,7 +1317,11 @@ def _chat_logprob_fingerprint_evidence(choice: Choice) -> dict[str, object] | No def values(items: Sequence[object] | None) -> list[dict[str, object]]: result: list[dict[str, object]] = [] for item in items or []: - data = _token_logprob_data(item) + data = ( + {**item.__dict__, **(item.model_extra or {})} + if type(item) is ChatCompletionTokenLogprob + else _token_logprob_data(item) + ) result.append( { key: data[key] diff --git a/tests/unit/trajectories/test_logprob_projection.py b/tests/unit/trajectories/test_logprob_projection.py new file mode 100644 index 000000000..95f0b7228 --- /dev/null +++ b/tests/unit/trajectories/test_logprob_projection.py @@ -0,0 +1,217 @@ +from collections import UserDict +from types import SimpleNamespace +from typing import Any, cast + +from openai.types.chat.chat_completion_token_logprob import ( + ChatCompletionTokenLogprob, + TopLogprob, +) +from pydantic import BaseModel +import pytest + +from art.trajectories._tokenize import ( + _chat_logprob_fingerprint_evidence, + _fingerprint, + _token_logprob_data, +) + +FIELDS = ("token", "token_id", "logprob", "bytes") + + +def evidence(row, refusal=False): + choice = SimpleNamespace( + logprobs=SimpleNamespace( + content=None if refusal else [row], refusal=[row] if refusal else None + ) + ) + result = _chat_logprob_fingerprint_evidence(cast(Any, choice)) + assert result is not None + return cast(list[dict[str, object]], result["refusal" if refusal else "content"])[0] + + +def reference(row): + data = _token_logprob_data(row) + return {key: data[key] for key in FIELDS if key in data} + + +@pytest.mark.parametrize("refusal", [False, True]) +@pytest.mark.parametrize( + "extra", [None, {}, {"token_id": 7}, {"token": None, "logprob": -0.0, "bytes": [9]}] +) +@pytest.mark.parametrize("missing", [None, "token", "logprob", "bytes"]) +def test_projection_preserves_order_precedence_missing_fields_and_aliases( + refusal, extra, missing +): + row = ChatCompletionTokenLogprob( + token="word", logprob=-0.2, bytes=[1], top_logprobs=[] + ) + object.__setattr__(row, "__pydantic_extra__", extra) + if missing is not None: + del row.__dict__[missing] + expected = reference(row) + actual = evidence(row, refusal) + assert list(actual) == list(expected) + assert _fingerprint(actual) == _fingerprint(expected) + assert all(actual[key] is value for key, value in expected.items()) + assert actual is not row.__dict__ and actual is not extra + + +@pytest.mark.parametrize("row_type", [dict, UserDict]) +def test_mapping_rows_keep_the_existing_projection(row_type): + row = row_type(token="word", token_id=7, logprob=-0.2, bytes=[1], ignored=[2]) + assert evidence(row) == reference(row) + assert evidence(row)["bytes"] is row["bytes"] + + +def test_typed_subclass_keeps_stored_field_and_extra_observation_order(): + calls = [] + + class Row(ChatCompletionTokenLogprob): + @property + def model_extra(self): + calls.append("extras") + self.__dict__["logprob"] = -9.0 + return {"token_id": 7} + + def model_dump(self, *args, **kwargs): + raise AssertionError("Typed row serialization must not execute") + + row = Row(token="word", logprob=-0.2, bytes=[1], top_logprobs=[]) + assert evidence(row) == { + "token": "word", + "token_id": 7, + "logprob": -0.2, + "bytes": [1], + } + assert row.logprob == -9.0 and calls == ["extras"] + + +def test_exact_row_captures_stored_fields_before_extra_property(monkeypatch): + row = ChatCompletionTokenLogprob( + token="word", logprob=-0.2, bytes=[1], top_logprobs=[] + ) + calls = [] + + def extras(self): + calls.append("extras") + self.__dict__["logprob"] = -9.0 + return {"token_id": 7} + + monkeypatch.setattr(ChatCompletionTokenLogprob, "model_extra", property(extras)) + assert evidence(row) == { + "token": "word", + "token_id": 7, + "logprob": -0.2, + "bytes": [1], + } + assert row.logprob == -9.0 and calls == ["extras"] + + +def test_rich_extra_mapping_keeps_unpacking_callbacks_and_precedence(): + calls = [] + + class Extras(UserDict): + def __getitem__(self, key): + calls.append(key) + return super().__getitem__(key) + + row = ChatCompletionTokenLogprob( + token="word", logprob=-0.2, bytes=[1], top_logprobs=[] + ) + object.__setattr__( + row, + "__pydantic_extra__", + Extras({"ignored": 1, "token_id": 7, "logprob": -0.4}), + ) + expected = reference(row) + observed = list(calls) + calls.clear() + assert evidence(row) == expected + assert calls == observed == ["ignored", "token_id", "logprob"] + + +def test_other_models_keep_serializer_output_and_exception_identity(): + failure = RuntimeError("row serialization failed") + calls = [] + + class Row(BaseModel): + def model_dump(self, *args, **kwargs): + calls.append(kwargs) + raise failure + + with pytest.raises(RuntimeError) as caught: + evidence(Row()) + assert caught.value is failure and calls == [{"mode": "python"}] + + +@pytest.mark.parametrize( + "field,value", + [("token", "changed"), ("token_id", 8), ("logprob", -9.0), ("bytes", [8])], +) +def test_projection_reads_fresh_mutable_evidence(field, value): + row = ChatCompletionTokenLogprob( + token="word", logprob=-0.2, bytes=[1], top_logprobs=[] + ) + before = _fingerprint(evidence(row)) + if field == "token_id": + cast(dict[str, Any], row.model_extra)[field] = value + else: + row.__dict__[field] = value + assert _fingerprint(evidence(row)) != before + + +@pytest.mark.parametrize("refusal", [False, True]) +def test_ignored_values_live_through_extra_observation(monkeypatch, refusal): + events = [] + held = [] + + class Alternative(TopLogprob): + def __del__(self): + events.append("finalizer") + held[0].__dict__["logprob"] = -9.0 + + row = ChatCompletionTokenLogprob( + token="word", + logprob=-0.2, + bytes=[1], + top_logprobs=[Alternative(token="word", logprob=-0.3, bytes=[2])], + ) + held.append(row) + + def extras(self): + events.append("extras-enter") + self.__dict__.pop("top_logprobs") + events.append("extras-read") + return {"logprob": self.__dict__["logprob"]} + + monkeypatch.setattr(ChatCompletionTokenLogprob, "model_extra", property(extras)) + assert evidence(row, refusal)["logprob"] == -0.2 + assert events == ["extras-enter", "extras-read", "finalizer"] + assert row.logprob == -9.0 + + +@pytest.mark.parametrize("refusal", [False, True]) +def test_rich_key_equality_follows_extra_observation(monkeypatch, refusal): + events = [] + + class Key(str): + __hash__ = str.__hash__ + + def __eq__(self, other): + events.append("key-equality") + return str.__eq__(self, other) + + row = ChatCompletionTokenLogprob( + token="word", logprob=-0.2, bytes=[1], top_logprobs=[] + ) + row.__dict__[Key("token_id")] = 7 + + def extras(self): + value = -9.0 if "key-equality" in events else -0.2 + events.append("extras-read") + return {"logprob": value} + + monkeypatch.setattr(ChatCompletionTokenLogprob, "model_extra", property(extras)) + result = evidence(row, refusal) + assert result["logprob"] == -0.2 and result["token_id"] == 7 + assert events == ["extras-read", "key-equality", "key-equality"] From 8468814844e09ff8d2fe09a351685003e8eea2b1 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 08:28:51 +0000 Subject: [PATCH 41/67] Preserve exact dict dispatch while narrowing its static type --- src/art/trajectories/_tokenize.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index 2690ceca2..4400e72ac 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -4874,7 +4874,7 @@ def snapshot(item: object) -> object: return kind, scalar, instance_state(item, getattr(item, "__dict__", None)) if kind in (list, tuple): result = kind, tuple(snapshot(child) for child in cast(Sequence, item)) - elif kind is dict or isinstance(item, Mapping): + elif type(item) is dict or isinstance(item, Mapping): result = ( kind, tuple((snapshot(key), snapshot(child)) for key, child in item.items()), From bec23b1fc7406f50a0ae279eaccc6c19106d26af Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 08:59:04 +0000 Subject: [PATCH 42/67] Narrow tokenizer regression fixture types for static checking --- .../trajectories/test_context_observation.py | 21 ++++++++++++++----- .../test_staged_tokenizer_guards.py | 4 +++- 2 files changed, 19 insertions(+), 6 deletions(-) diff --git a/tests/unit/trajectories/test_context_observation.py b/tests/unit/trajectories/test_context_observation.py index e70d94ece..79c846cdb 100644 --- a/tests/unit/trajectories/test_context_observation.py +++ b/tests/unit/trajectories/test_context_observation.py @@ -1,4 +1,5 @@ from collections import UserDict +from typing import Any, cast import pytest @@ -6,6 +7,8 @@ class ObservedDict(dict): + reads: int + def items(self): self.reads += 1 return super().items() @@ -17,7 +20,9 @@ def test_context_preserves_mapping_order_aliases_and_fresh_mutable_reads(mapping mapping = mapping_type({"first": child, "second": child}) if isinstance(mapping, ObservedDict): mapping.reads = 0 - original = _tokenization_context([mapping, mapping]) + original = cast( + tuple[type, tuple[Any, Any]], _tokenization_context([mapping, mapping]) + ) first, second = original[1] assert first is second assert first[0] is mapping_type @@ -34,14 +39,20 @@ def test_context_preserves_mapping_order_aliases_and_fresh_mutable_reads(mapping def test_context_rejects_stale_identity_entry_without_borrowing_its_value(): value = {"request": [1, True, "1"]} - observed = {id(value): (object(), "stale")} - assert _tokenization_context(value, _observed=observed) == _tokenization_context(value) + observed: dict[int, tuple[object, object]] = {id(value): (object(), "stale")} + assert _tokenization_context(value, _observed=observed) == _tokenization_context( + value + ) assert observed[id(value)][0] is value -@pytest.mark.parametrize("left,right", [(1, True), (0.0, -0.0), ("x", b"x"), (None, "None")]) +@pytest.mark.parametrize( + "left,right", [(1, True), (0.0, -0.0), ("x", b"x"), (None, "None")] +) def test_context_retains_typed_scalar_distinctions(left, right): - assert _tokenization_context({"value": left}) != _tokenization_context({"value": right}) + assert _tokenization_context({"value": left}) != _tokenization_context( + {"value": right} + ) def test_context_preserves_custom_mapping_exception_identity(): diff --git a/tests/unit/trajectories/test_staged_tokenizer_guards.py b/tests/unit/trajectories/test_staged_tokenizer_guards.py index 326e8001b..966126b5c 100644 --- a/tests/unit/trajectories/test_staged_tokenizer_guards.py +++ b/tests/unit/trajectories/test_staged_tokenizer_guards.py @@ -79,7 +79,9 @@ def observe(self, *args, **kwargs): monkeypatch.setattr(_tokenize._ChatViewTokenizer, "__init__", observe) exchange = _chat_exchange([1], [2, 9]) - exchange.response.choices[0].model_extra.pop("prompt_token_ids") + extra = exchange.response.choices[0].model_extra + assert extra is not None + extra.pop("prompt_token_ids") history = tr.Trajectory( exchanges=tr.TrajectoryExchanges(chat_completions=[exchange]) ).chat_completions_history() From da8106f5968ed6327d3c04619db7df4cd851d4ca Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 09:12:39 +0000 Subject: [PATCH 43/67] Validate physical mapping state across tokenizer callbacks --- src/art/trajectories/_tokenize.py | 16 ++- .../trajectories/test_context_observation.py | 11 +- .../test_mapping_context_state.py | 120 ++++++++++++++++++ 3 files changed, 139 insertions(+), 8 deletions(-) create mode 100644 tests/unit/trajectories/test_mapping_context_state.py diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index 926ccca6a..45d44ef2a 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -17,7 +17,7 @@ from pickle import Pickler, PicklingError import re import threading -from types import FunctionType, MemberDescriptorType +from types import FunctionType, GetSetDescriptorType, MemberDescriptorType from typing import TYPE_CHECKING, Any, Literal, Protocol, cast import warnings @@ -4879,6 +4879,18 @@ def snapshot(item: object) -> object: kind, tuple((snapshot(key), snapshot(child)) for key, child in item.items()), ) + if kind is not dict: + # A renderer can read mapping attributes as well as its items. + # Inspect physical instance storage without invoking overrides. + dictionary = None + for owner in type.__getattribute__(kind, "__mro__"): + descriptor = type.__getattribute__(owner, "__dict__").get( + "__dict__" + ) + if type(descriptor) is GetSetDescriptorType: + dictionary = descriptor.__get__(item, kind) + break + result = (*result, instance_state(item, dictionary)) elif isinstance(item, Exchange): result = kind, identity, item.model, snapshot(item.request) elif isinstance(item, BaseModel): @@ -4896,7 +4908,7 @@ def snapshot(item: object) -> object: return result def instance_state(item: object, dictionary: object) -> object: - # Scalar subclasses and Enum members may keep mutable state in slots. + # Rich values may keep mutable state in inherited or shadowed slots. # Read actual slot storage, including inherited/shadowed slots, without # invoking an instance's attribute lookup or replacement properties. slots = [] diff --git a/tests/unit/trajectories/test_context_observation.py b/tests/unit/trajectories/test_context_observation.py index 79c846cdb..585e6de60 100644 --- a/tests/unit/trajectories/test_context_observation.py +++ b/tests/unit/trajectories/test_context_observation.py @@ -5,12 +5,12 @@ from art.trajectories._tokenize import _tokenization_context +_observed_mappings: list[object] = [] -class ObservedDict(dict): - reads: int +class ObservedDict(dict): def items(self): - self.reads += 1 + _observed_mappings.append(self) return super().items() @@ -18,8 +18,7 @@ def items(self): def test_context_preserves_mapping_order_aliases_and_fresh_mutable_reads(mapping_type): child = ["before"] mapping = mapping_type({"first": child, "second": child}) - if isinstance(mapping, ObservedDict): - mapping.reads = 0 + _observed_mappings.clear() original = cast( tuple[type, tuple[Any, Any]], _tokenization_context([mapping, mapping]) ) @@ -29,7 +28,7 @@ def test_context_preserves_mapping_order_aliases_and_fresh_mutable_reads(mapping assert [entry[0] for entry in first[1]] == [(str, "first"), (str, "second")] assert first[1][0][1] is first[1][1][1] if isinstance(mapping, ObservedDict): - assert mapping.reads == 1 + assert _observed_mappings == [mapping] child[0] = "after" assert _tokenization_context([mapping, mapping]) != original child[0] = "before" diff --git a/tests/unit/trajectories/test_mapping_context_state.py b/tests/unit/trajectories/test_mapping_context_state.py new file mode 100644 index 000000000..9054a75fc --- /dev/null +++ b/tests/unit/trajectories/test_mapping_context_state.py @@ -0,0 +1,120 @@ +from __future__ import annotations + +from collections import UserDict +from typing import Any, cast + +import pytest +from test_tokenize import _chat_exchange + +import art.trajectories as tr +from art.trajectories import _tokenize + + +class _Dictionary(dict): + revision: int + + +class _UserDictionary(UserDict): + revision: int + + +class _Slots(dict): + __slots__ = ("revision",) + revision: int + + +class _ShadowSlots(_Slots): + __slots__ = ("revision",) + + +class _HiddenDictionary(_Dictionary): + @property + def __dict__(self): + raise AssertionError("context snapshots must inspect physical instance storage") + + +def _options(kind): + types = { + "dict": _Dictionary, + "user_dict": _UserDictionary, + "slots": _Slots, + "shadowed_slot": _ShadowSlots, + "hidden_dict": _HiddenDictionary, + } + options = types[kind](key="stable") + options.revision = 0 + if kind == "shadowed_slot": + descriptor = _Slots.__dict__["revision"] + descriptor.__set__(options, 0) + return ( + options, + lambda: descriptor.__get__(options), + lambda value: descriptor.__set__(options, value), + ) + return ( + options, + lambda: options.revision, + lambda value: setattr(options, "revision", value), + ) + + +@pytest.mark.parametrize( + "kind", ["dict", "user_dict", "slots", "shadowed_slot", "hidden_dict"] +) +def test_mapping_snapshot_observes_physical_attributes_and_slots(kind): + options, read, write = _options(kind) + original_items = list(options.items()) + before = _tokenize._tokenization_context([options, options]) + write(1) + assert read() == 1 and list(options.items()) == original_items + assert _tokenize._tokenization_context([options, options]) != before + write(0) + assert _tokenize._tokenization_context([options, options]) == before + + +@pytest.mark.parametrize( + "kind", ["dict", "user_dict", "slots", "shadowed_slot", "hidden_dict"] +) +@pytest.mark.parametrize("mutate", [False, True]) +def test_mapping_attribute_callback_cannot_change_rendering_or_restore_later( + kind, mutate +): + options, read, write = _options(kind) + exchange = _chat_exchange([1], [2]) + exchange.request["messages"] = [{"role": "user", "content": "q"}] + value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(chat_completions=[exchange])) + before = value.model_dump_json() + calls = [] + + class Tokenizer: + def apply_chat_template(self, messages, **kwargs): + assert kwargs["options"] is options + calls.append("render") + if mutate: + write(1) + return [3 if read() else 1, 2] + + def __call__(self, text, **kwargs): + calls.append("encode") + # A later callback cannot hide the renderer's mutation. + if mutate: + write(0) + return [2 if text == "answer" else 1] + + def invoke(): + return value.tokenize( + tokenizer=cast(Any, Tokenizer()), + chat_template="custom", + chat_template_kwargs={"options": options}, + ) + + if mutate: + with pytest.raises(ValueError, match="[Cc]ontext changed"): + invoke() + assert calls == ["render"] and read() == 1 + else: + actual = invoke() + assert actual.tokens == [1, 2] and actual.logprobs[-1] == -0.2 + assert actual.flags[-1] & tr.TokenFlag.SAMPLED + assert read() == 0 and "render" in calls + assert value.model_dump_json() == before From 6878aa38c2b26ba8515a79482ae9b0ef5c26a2fe Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 09:28:22 +0000 Subject: [PATCH 44/67] Treat recursive tokenization contexts as unsupported callback proofs --- src/art/trajectories/_tokenize.py | 7 +- .../trajectories/test_recursive_context.py | 119 ++++++++++++++++++ 2 files changed, 125 insertions(+), 1 deletion(-) create mode 100644 tests/unit/trajectories/test_recursive_context.py diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index 45d44ef2a..1328482cc 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -4924,7 +4924,12 @@ def instance_state(item: object, dictionary: object) -> object: slots.append((owner, name, True, snapshot(child))) return snapshot(dictionary), tuple(slots) - return snapshot(value) + try: + return snapshot(value) + except RecursionError as error: + # Recursive context cannot prove callback stability, but complete native + # records can still use the ordinary opaque-context bypass. + raise TypeError("Unsupported recursive tokenization context") from error def _tokenization_context_validator(value: object) -> Callable[[bool], None]: diff --git a/tests/unit/trajectories/test_recursive_context.py b/tests/unit/trajectories/test_recursive_context.py new file mode 100644 index 000000000..f4fc3ba98 --- /dev/null +++ b/tests/unit/trajectories/test_recursive_context.py @@ -0,0 +1,119 @@ +from typing import Any, cast + +import pytest +from test_tokenize import ( + _chat_exchange, + _completion_exchange, + _message_exchange, + _response_exchange, +) + +import art.trajectories as tr +from art.trajectories._tokenize import _tokenization_context + + +class _Options(dict): + owner: object + + +class _Slots(dict): + __slots__ = ("owner",) + owner: object + + +def _recursive(kind): + if kind == "list": + value = [] + value.append(value) + return value + value = _Options(key="stable") if kind == "dictionary" else _Slots(key="stable") + value.owner = value + return value + + +@pytest.mark.parametrize("kind", ["dictionary", "slots", "list"]) +def test_recursive_context_has_an_unsupported_proof_boundary(kind): + value = _recursive(kind) + with pytest.raises(TypeError, match="Unsupported recursive tokenization context"): + _tokenization_context(value) + + +@pytest.mark.parametrize("kind", ["dictionary", "slots", "list"]) +@pytest.mark.parametrize("authority", ["none", "plain", "dynamic"]) +@pytest.mark.parametrize( + "protocol", ["chat_completions", "messages", "responses", "completions"] +) +def test_native_recursive_context_bypasses_without_callbacks( + kind, authority, protocol, monkeypatch +): + calls = [] + + class PlainTokenizer: + eos_token_id = 2 + + class DynamicTokenizer: + @property + def eos_token_id(self): + calls.append("eos") + return 2 + + if protocol == "chat_completions": + exchange = _chat_exchange([1], [2]) + elif protocol == "messages": + exchange = _message_exchange( + cast(Any, dict(model="test/model", max_tokens=5, messages=[])), + prompt_token_ids=[1], + token_ids=[2], + logprobs=[-0.2], + ) + elif protocol == "responses": + exchange = _response_exchange("public-recursive", 2, prompt_token_ids=[1]) + else: + exchange = _completion_exchange(prompt="question") + context = _recursive(kind) + exchange.request["metadata"] = cast(Any, {"unused": context}) + value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(**{protocol: [exchange]})) + + def unexpected(*args, **kwargs): + pytest.fail("Complete native records must not load or render") + + monkeypatch.setattr("art.trajectories._tokenize._load_tokenizer", unexpected) + tokenizer = ( + None + if authority == "none" + else PlainTokenizer() + if authority == "plain" + else DynamicTokenizer() + ) + expected_calls = [] + if authority == "dynamic": + # Preserve the existing opaque-context callback/refusal boundary. + exchange.request["metadata"] = cast(Any, {"unused": object()}) + with pytest.raises(ValueError, match="[Cc]ontext"): + value.tokenize(tokenizer=cast(Any, tokenizer)) + expected_calls = calls.copy() + calls.clear() + exchange.request["metadata"] = cast(Any, {"unused": context}) + with pytest.raises(ValueError, match="[Cc]ontext"): + value.tokenize(tokenizer=cast(Any, tokenizer)) + else: + result = value.tokenize(tokenizer=cast(Any, tokenizer)) + assert result.tokens == [1, 2] + assert result.flags[-1] & tr.TokenFlag.SAMPLED + assert bool(result.flags[-1] & tr.TokenFlag.STOP) == (authority == "plain") + assert calls == expected_calls + assert cast(Any, exchange.request["metadata"])["unused"] is context + + +def test_actual_renderer_recursion_error_is_not_context_admission(): + exchange = _chat_exchange([1], [2]) + value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(chat_completions=[exchange])) + failure = RecursionError("renderer failure") + + class Tokenizer: + def apply_chat_template(self, messages, **kwargs): + raise failure + + with pytest.raises(RecursionError) as actual: + value.tokenize(tokenizer=cast(Any, Tokenizer()), chat_template="custom") + assert actual.value is failure From 20341984753d406e1a73cacee4df5cce418e370f Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 09:32:44 +0000 Subject: [PATCH 45/67] Require unsupported-context error while retaining recursion cause --- tests/unit/trajectories/test_resolved_stop_authority.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/tests/unit/trajectories/test_resolved_stop_authority.py b/tests/unit/trajectories/test_resolved_stop_authority.py index e22a9c3f6..4b45af2a3 100644 --- a/tests/unit/trajectories/test_resolved_stop_authority.py +++ b/tests/unit/trajectories/test_resolved_stop_authority.py @@ -580,8 +580,9 @@ def test_context_snapshot_shared_values_are_fresh_between_observations(): def test_context_snapshot_does_not_certify_cyclic_or_opaque_values(): cyclic = [] cyclic.append(cyclic) - with pytest.raises(RecursionError): + with pytest.raises(TypeError, match="Unsupported recursive") as failure: module._tokenization_context(cyclic) + assert isinstance(failure.value.__cause__, RecursionError) with pytest.raises(TypeError, match="Unsupported mutable"): module._tokenization_context([object()]) From dd5d7455e5122896303b2f996b9f3b801bdd35eb Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 09:41:20 +0000 Subject: [PATCH 46/67] Observe complete physical Pydantic context state across callbacks --- src/art/trajectories/_tokenize.py | 23 ++-- .../trajectories/test_model_context_state.py | 120 ++++++++++++++++++ 2 files changed, 132 insertions(+), 11 deletions(-) create mode 100644 tests/unit/trajectories/test_model_context_state.py diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index 1328482cc..43b4f1835 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -4872,8 +4872,8 @@ def snapshot(item: object) -> object: else: scalar = bytes.__bytes__(item) return kind, scalar, instance_state(item, getattr(item, "__dict__", None)) - if kind in (list, tuple): - result = kind, tuple(snapshot(child) for child in cast(Sequence, item)) + if kind in (list, tuple, set, frozenset): + result = kind, tuple(snapshot(child) for child in cast(Iterable, item)) elif type(item) is dict or isinstance(item, Mapping): result = ( kind, @@ -4882,15 +4882,7 @@ def snapshot(item: object) -> object: if kind is not dict: # A renderer can read mapping attributes as well as its items. # Inspect physical instance storage without invoking overrides. - dictionary = None - for owner in type.__getattribute__(kind, "__mro__"): - descriptor = type.__getattribute__(owner, "__dict__").get( - "__dict__" - ) - if type(descriptor) is GetSetDescriptorType: - dictionary = descriptor.__get__(item, kind) - break - result = (*result, instance_state(item, dictionary)) + result = (*result, instance_state(item, instance_dictionary(item))) elif isinstance(item, Exchange): result = kind, identity, item.model, snapshot(item.request) elif isinstance(item, BaseModel): @@ -4901,12 +4893,21 @@ def snapshot(item: object) -> object: for name in type(item).model_fields ), snapshot(item.model_extra), + instance_state(item, instance_dictionary(item)), ) else: raise TypeError("Unsupported mutable tokenization context") observed[identity] = item, result return result + def instance_dictionary(item: object) -> object: + kind = type(item) + for owner in type.__getattribute__(kind, "__mro__"): + descriptor = type.__getattribute__(owner, "__dict__").get("__dict__") + if type(descriptor) is GetSetDescriptorType: + return descriptor.__get__(item, kind) + return None + def instance_state(item: object, dictionary: object) -> object: # Rich values may keep mutable state in inherited or shadowed slots. # Read actual slot storage, including inherited/shadowed slots, without diff --git a/tests/unit/trajectories/test_model_context_state.py b/tests/unit/trajectories/test_model_context_state.py new file mode 100644 index 000000000..4b4d75859 --- /dev/null +++ b/tests/unit/trajectories/test_model_context_state.py @@ -0,0 +1,120 @@ +from functools import cached_property +from typing import Any, cast + +from pydantic import BaseModel, PrivateAttr +import pytest +from test_tokenize import _chat_exchange + +import art.trajectories as tr +from art.trajectories import _tokenize + + +class _Model(BaseModel): + value: int = 0 + _revision: list[int] = PrivateAttr(default_factory=lambda: [0]) + + @cached_property + def cached(self): + return 0 + + +class _Slots(_Model): + __slots__ = ("revision",) + + +class _HiddenDictionary(_Model): + @property + def __dict__(self): + raise AssertionError("model context must inspect physical instance storage") + + @__dict__.setter + def __dict__(self, value): + BaseModel.__dict__["__dict__"].__set__(self, value) + + +def _options(kind): + value = _HiddenDictionary() if kind == "hidden_dict" else _Model() + if kind in ("private", "hidden_dict"): + return ( + value, + lambda: value._revision[0], + lambda n: value._revision.__setitem__(0, n), + ) + if kind == "fields_set": + + def write(n): + if n: + value.value = 0 + else: + value.model_fields_set.clear() + + return value, lambda: len(value.model_dump(exclude_unset=True)), write + if kind == "cached": + assert value.cached == 0 + return value, lambda: value.cached, lambda n: setattr(value, "cached", n) + value = _Slots() + object.__setattr__(value, "revision", 0) + descriptor = _Slots.__dict__["revision"] + return ( + value, + lambda: descriptor.__get__(value), + lambda n: descriptor.__set__(value, n), + ) + + +@pytest.mark.parametrize( + "kind", ["private", "fields_set", "cached", "slots", "hidden_dict"] +) +def test_model_context_observes_physical_state_and_restoration(kind): + options, read, write = _options(kind) + before = _tokenize._tokenization_context([options, options]) + write(1) + assert read() == 1 + assert _tokenize._tokenization_context([options, options]) != before + write(0) + assert _tokenize._tokenization_context([options, options]) == before + + +@pytest.mark.parametrize( + "kind", ["private", "fields_set", "cached", "slots", "hidden_dict"] +) +@pytest.mark.parametrize("mutate", [False, True]) +def test_model_state_callback_cannot_change_rendering_or_restore_later(kind, mutate): + options, read, write = _options(kind) + exchange = _chat_exchange([1], [2]) + exchange.request["messages"] = [{"role": "user", "content": "q"}] + value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(chat_completions=[exchange])) + before = value.model_dump_json() + calls = [] + + class Tokenizer: + def apply_chat_template(self, messages, **kwargs): + assert kwargs["options"] is options + calls.append("render") + if mutate: + write(1) + return [3 if read() else 1, 2] + + def __call__(self, text, **kwargs): + calls.append("encode") + if mutate: + write(0) + return [2 if text == "answer" else 1] + + def invoke(): + return value.tokenize( + tokenizer=cast(Any, Tokenizer()), + chat_template="custom", + chat_template_kwargs={"options": options}, + ) + + if mutate: + with pytest.raises(ValueError, match="[Cc]ontext changed"): + invoke() + assert calls == ["render"] and read() == 1 + else: + actual = invoke() + assert actual.tokens == [1, 2] and actual.logprobs[-1] == -0.2 + assert actual.flags[-1] & tr.TokenFlag.SAMPLED + assert read() == 0 and "render" in calls + assert value.model_dump_json() == before From 7951618d6961678438dcc3598d9129b5ce8b907f Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 09:40:02 +0000 Subject: [PATCH 47/67] Release tokenization snapshot closures after each observation --- src/art/trajectories/_tokenize.py | 3 + .../test_context_snapshot_lifetime.py | 96 +++++++++++++++++++ 2 files changed, 99 insertions(+) create mode 100644 tests/unit/trajectories/test_context_snapshot_lifetime.py diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index 43b4f1835..5ebf80e12 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -4931,6 +4931,9 @@ def instance_state(item: object, dictionary: object) -> object: # Recursive context cannot prove callback stability, but complete native # records can still use the ordinary opaque-context bypass. raise TypeError("Unsupported recursive tokenization context") from error + finally: + # Release the recursive closures and their observation memo promptly. + del snapshot, instance_state def _tokenization_context_validator(value: object) -> Callable[[bool], None]: diff --git a/tests/unit/trajectories/test_context_snapshot_lifetime.py b/tests/unit/trajectories/test_context_snapshot_lifetime.py new file mode 100644 index 000000000..d174390e3 --- /dev/null +++ b/tests/unit/trajectories/test_context_snapshot_lifetime.py @@ -0,0 +1,96 @@ +from collections import UserDict +import gc +from typing import cast +import weakref + +import pytest + +from art.trajectories._tokenize import ( + _tokenization_context, + _tokenization_context_validator, +) + + +class _Options(dict): + opaque: object + + +class _Slots(dict): + __slots__ = ("revision", "__weakref__") + + +@pytest.fixture +def no_cyclic_gc(): + enabled = gc.isenabled() + gc.disable() + try: + yield + finally: + if enabled: + gc.enable() + gc.collect() + + +@pytest.mark.parametrize("kind", [_Options, _Slots, UserDict]) +def test_completed_snapshot_releases_input_without_cyclic_gc(kind, no_cyclic_gc): + references = [] + snapshots = [] + for _ in range(32): + value = kind(key=["stable"]) + value.revision = ["before"] + references.append(weakref.ref(value)) + snapshots.append( + cast( + tuple[type, tuple[object, object]], + _tokenization_context([value, value]), + ) + ) + del value + assert all(reference() is None for reference in references) + assert all(value[1][0] is value[1][1] for value in snapshots) + assert all(value == snapshots[0] for value in snapshots) + + +def test_caller_memo_retains_ownership_until_caller_clears_it(no_cyclic_gc): + value = _Options(key=["stable"]) + reference = weakref.ref(value) + observed = {} + snapshot = _tokenization_context(value, _observed=observed) + identity = id(value) + assert observed[identity] == (value, snapshot) + del value + assert reference() is not None + assert observed[identity][1] is snapshot + observed.clear() + assert reference() is None + assert snapshot[1][0][0] == (str, "key") + + +@pytest.mark.parametrize("partial", [False, True]) +def test_failed_snapshot_releases_input_after_error_leaves_scope(partial, no_cyclic_gc): + value = _Options(key=["stable"]) + if not partial: + value.opaque = object() + reference = weakref.ref(value) + + def fail(): + try: + _tokenization_context([value, object()] if partial else value) + except TypeError: + return + pytest.fail("Opaque context unexpectedly accepted") + + fail() + del value + assert reference() is None + + +def test_validator_owns_input_only_until_validator_is_released(no_cyclic_gc): + value = _Options(key=["stable"]) + reference = weakref.ref(value) + validate = _tokenization_context_validator(value) + validate(True) + del value + assert reference() is not None + del validate + assert reference() is None From e32455a8c55eab0ca5c064cabe1e7ac672b5679f Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 10:04:41 +0000 Subject: [PATCH 48/67] Detect recursive context explicitly without hiding observer errors --- src/art/trajectories/_tokenize.py | 16 +++-- .../test_context_observer_errors.py | 72 +++++++++++++++++++ .../test_resolved_stop_authority.py | 3 +- 3 files changed, 84 insertions(+), 7 deletions(-) create mode 100644 tests/unit/trajectories/test_context_observer_errors.py diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index 5ebf80e12..eb56142c4 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -4840,6 +4840,7 @@ def _tokenization_context( # Shared message dictionaries occur in many recorded request prefixes. # Intern only within this one observation, never across callback checks. observed: dict[int, tuple[object, object]] = {} if _observed is None else _observed + active: set[int] = set() def snapshot(item: object) -> object: kind = type(item) @@ -4851,6 +4852,15 @@ def snapshot(item: object) -> object: result = kind, repr(item) if kind is float else item observed[identity] = item, result return result + if identity in active: + raise TypeError("Unsupported recursive tokenization context") + active.add(identity) + try: + return snapshot_compound(item, kind, identity) + finally: + active.remove(identity) + + def snapshot_compound(item: object, kind: type, identity: int) -> object: if isinstance(item, Enum): if getattr(item, "__objclass__", kind) is not kind: raise TypeError("Unsupported enum tokenization context") @@ -4927,13 +4937,9 @@ def instance_state(item: object, dictionary: object) -> object: try: return snapshot(value) - except RecursionError as error: - # Recursive context cannot prove callback stability, but complete native - # records can still use the ordinary opaque-context bypass. - raise TypeError("Unsupported recursive tokenization context") from error finally: # Release the recursive closures and their observation memo promptly. - del snapshot, instance_state + del snapshot, snapshot_compound, instance_state def _tokenization_context_validator(value: object) -> Callable[[bool], None]: diff --git a/tests/unit/trajectories/test_context_observer_errors.py b/tests/unit/trajectories/test_context_observer_errors.py new file mode 100644 index 000000000..d1e2bb95c --- /dev/null +++ b/tests/unit/trajectories/test_context_observer_errors.py @@ -0,0 +1,72 @@ +from typing import Any, cast + +from pydantic import BaseModel +import pytest +from test_tokenize import _chat_exchange + +import art.trajectories as tr +from art.trajectories._tokenize import _tokenization_context + + +def _observer(kind, failure): + if kind in ("items", "partial_items"): + + class Options(dict): + def items(self): + if kind == "partial_items": + yield "stable", [1] + raise failure + + return Options(key="stable") + + class Model(BaseModel): + value: str = "stable" + + def __getattribute__(self, name): + if name == kind: + raise failure + return super().__getattribute__(name) + + return Model() + + +@pytest.mark.parametrize("kind", ["items", "partial_items", "value", "model_extra"]) +@pytest.mark.parametrize("error_type", [RecursionError, RuntimeError]) +@pytest.mark.parametrize("route", ["snapshot", "native", "render"]) +def test_observer_failure_is_not_an_unsupported_context(kind, error_type, route): + failure = error_type("observer failure") + options = _observer(kind, failure) + calls = [] + + class Tokenizer: + def apply_chat_template(self, messages, **kwargs): + calls.append("render") + return [1, 2] + + def __call__(self, text, **kwargs): + calls.append("encode") + return [2 if text == "answer" else 1] + + def invoke(): + if route == "snapshot": + return _tokenization_context(options) + exchange = _chat_exchange([1], [2]) + if route == "native": + exchange.request["metadata"] = cast(Any, {"unused": options}) + value = tr.Trajectory( + exchanges=tr.TrajectoryExchanges(chat_completions=[exchange]) + ) + return ( + value.tokenize() + if route == "native" + else value.tokenize( + tokenizer=cast(Any, Tokenizer()), + chat_template="custom", + chat_template_kwargs={"options": options}, + ) + ) + + with pytest.raises(error_type) as actual: + invoke() + assert actual.value is failure + assert calls == [] diff --git a/tests/unit/trajectories/test_resolved_stop_authority.py b/tests/unit/trajectories/test_resolved_stop_authority.py index 4b45af2a3..c50336564 100644 --- a/tests/unit/trajectories/test_resolved_stop_authority.py +++ b/tests/unit/trajectories/test_resolved_stop_authority.py @@ -580,9 +580,8 @@ def test_context_snapshot_shared_values_are_fresh_between_observations(): def test_context_snapshot_does_not_certify_cyclic_or_opaque_values(): cyclic = [] cyclic.append(cyclic) - with pytest.raises(TypeError, match="Unsupported recursive") as failure: + with pytest.raises(TypeError, match="Unsupported recursive"): module._tokenization_context(cyclic) - assert isinstance(failure.value.__cause__, RecursionError) with pytest.raises(TypeError, match="Unsupported mutable"): module._tokenization_context([object()]) From 9b51a1956156f47da7984e00c9594431bf63e2b4 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 10:18:09 +0000 Subject: [PATCH 49/67] Observe native instance storage consistently across context types --- src/art/trajectories/_tokenize.py | 14 +- .../trajectories/test_scalar_context_state.py | 214 ++++++++++++++++++ 2 files changed, 222 insertions(+), 6 deletions(-) create mode 100644 tests/unit/trajectories/test_scalar_context_state.py diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index eb56142c4..cf1fda964 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -4868,7 +4868,7 @@ def snapshot_compound(item: object, kind: type, identity: int) -> object: item, { key: child - for key, child in vars(item).items() + for key, child in cast(dict, instance_dictionary(item)).items() if key != "__objclass__" }, ) @@ -4881,7 +4881,7 @@ def snapshot_compound(item: object, kind: type, identity: int) -> object: scalar = repr(float.__float__(item)) else: scalar = bytes.__bytes__(item) - return kind, scalar, instance_state(item, getattr(item, "__dict__", None)) + return kind, scalar, instance_state(item, instance_dictionary(item)) if kind in (list, tuple, set, frozenset): result = kind, tuple(snapshot(child) for child in cast(Iterable, item)) elif type(item) is dict or isinstance(item, Mapping): @@ -4912,10 +4912,12 @@ def snapshot_compound(item: object, kind: type, identity: int) -> object: def instance_dictionary(item: object) -> object: kind = type(item) - for owner in type.__getattribute__(kind, "__mro__"): - descriptor = type.__getattribute__(owner, "__dict__").get("__dict__") + for owner in type.__dict__["__mro__"].__get__(kind): + descriptor = type.__dict__["__dict__"].__get__(owner).get("__dict__") if type(descriptor) is GetSetDescriptorType: return descriptor.__get__(item, kind) + if type.__dict__["__dictoffset__"].__get__(kind): + raise TypeError("Unsupported hidden instance dictionary") return None def instance_state(item: object, dictionary: object) -> object: @@ -4924,8 +4926,8 @@ def instance_state(item: object, dictionary: object) -> object: # invoking an instance's attribute lookup or replacement properties. slots = [] kind = type(item) - for owner in type.__getattribute__(kind, "__mro__"): - for name, descriptor in type.__getattribute__(owner, "__dict__").items(): + for owner in type.__dict__["__mro__"].__get__(kind): + for name, descriptor in type.__dict__["__dict__"].__get__(owner).items(): if type(descriptor) is MemberDescriptorType: try: child = descriptor.__get__(item, kind) diff --git a/tests/unit/trajectories/test_scalar_context_state.py b/tests/unit/trajectories/test_scalar_context_state.py new file mode 100644 index 000000000..bdf8d4ca7 --- /dev/null +++ b/tests/unit/trajectories/test_scalar_context_state.py @@ -0,0 +1,214 @@ +from enum import Enum +from typing import Any, cast + +import pytest +from test_tokenize import _chat_exchange + +import art.trajectories as tr +from art.trajectories import _tokenize + + +def _scalar(kind, override): + accesses = [] + + def dictionary(self): + accesses.append("dictionary") + if override == "raising": + raise AssertionError("physical storage must bypass the override") + return {} + + if kind in ("enum", "string_enum"): + + class Options(Enum): + value = "stable" + + __dict__ = property(dictionary) + + class TextOptions(str, Enum): + value = "stable" + + __dict__ = property(dictionary) + + options = TextOptions.value if kind == "string_enum" else Options.value + else: + base, value = { + "str": (str, "stable"), + "int": (int, 7), + "float": (float, 0.5), + "bytes": (bytes, b"stable"), + }[kind] + physical = type("Physical", (base,), {}) + hidden = type("Hidden", (physical,), {"__dict__": property(dictionary)}) + options = hidden(value) + setattr(options, "revision", [0]) + return options, accesses + + +@pytest.mark.parametrize( + "kind", ["str", "int", "float", "bytes", "enum", "string_enum"] +) +@pytest.mark.parametrize("override", ["hidden", "raising"]) +def test_scalar_snapshot_tracks_physical_state_and_restoration(kind, override): + options, accesses = _scalar(kind, override) + context = [options, options] + before = _tokenize._tokenization_context(context) + validate = _tokenize._tokenization_context_validator(context) + options.revision[0] = 1 + assert _tokenize._tokenization_context(context) != before + with pytest.raises(ValueError, match="context changed"): + validate(True) + options.revision[0] = 0 + assert _tokenize._tokenization_context(context) == before + validate(True) + assert accesses == [] + + +@pytest.mark.parametrize( + "kind", ["str", "int", "float", "bytes", "enum", "string_enum"] +) +@pytest.mark.parametrize("override", ["hidden", "raising"]) +@pytest.mark.parametrize("mutate", [False, True]) +def test_scalar_callback_cannot_hide_mutation_or_restore_later(kind, override, mutate): + options, accesses = _scalar(kind, override) + _check_callback(options, accesses, mutate) + + +def _check_callback(options, accesses, mutate): + exchange = _chat_exchange([1], [2]) + exchange.request["messages"] = [{"role": "user", "content": "q"}] + value = tr.Trajectory(exchanges=tr.TrajectoryExchanges(chat_completions=[exchange])) + before = value.model_dump_json() + calls = [] + + class Tokenizer: + def apply_chat_template(self, messages, **kwargs): + assert kwargs["options"] is options + calls.append("render") + if mutate: + options.revision[0] = 1 + return [3 if options.revision[0] else 1, 2] + + def __call__(self, text, **kwargs): + calls.append("encode") + if mutate: + options.revision[0] = 0 + return [2 if text == "answer" else 1] + + def invoke(): + return value.tokenize( + tokenizer=cast(Any, Tokenizer()), + chat_template="custom", + chat_template_kwargs={"options": options}, + ) + + if mutate: + with pytest.raises(ValueError, match="[Cc]ontext changed"): + invoke() + assert calls == ["render"] and options.revision == [1] + else: + actual = invoke() + assert actual.tokens == [1, 2] and actual.logprobs[-1] == -0.2 + assert actual.flags[-1] & tr.TokenFlag.SAMPLED + assert options.revision == [0] and "render" in calls + assert accesses == [] + assert value.model_dump_json() == before + + +@pytest.mark.parametrize( + "base,value", [(str, "x"), (int, 1), (float, 0.5), (bytes, b"x"), (dict, {})] +) +@pytest.mark.parametrize("override", ["hidden", "raising"]) +@pytest.mark.parametrize("route", ["snapshot", "native", "render"]) +def test_direct_hidden_dictionary_cannot_be_certified(base, value, override, route): + accesses = [] + + def dictionary(self): + accesses.append("dictionary") + if override == "raising": + raise AssertionError("dictionary override must not be called") + return {} + + direct = type("Direct", (base,), {"__dict__": property(dictionary)}) + options = direct(value) + options.revision = [0] + calls = [] + + class Tokenizer: + def apply_chat_template(self, messages, **kwargs): + calls.append("render") + options.revision[0] = 1 + return [1, 2] + + def __call__(self, text, **kwargs): + calls.append("encode") + return [2 if text == "answer" else 1] + + if route == "snapshot": + with pytest.raises(TypeError, match="Unsupported hidden instance dictionary"): + _tokenize._tokenization_context(options) + validate = _tokenize._tokenization_context_validator(options) + validate(False) + with pytest.raises(ValueError, match="cannot be checked"): + validate(True) + else: + exchange = _chat_exchange([1], [2]) + trajectory = tr.Trajectory( + exchanges=tr.TrajectoryExchanges(chat_completions=[exchange]) + ) + if route == "native": + exchange.request["metadata"] = cast(Any, {"unused": options}) + assert trajectory.tokenize().tokens == [1, 2] + else: + with pytest.raises( + TypeError, match="Unsupported hidden instance dictionary" + ): + trajectory.tokenize( + tokenizer=cast(Any, Tokenizer()), + chat_template="custom", + chat_template_kwargs={"options": options}, + ) + assert accesses == calls == [] + assert options.revision == [0] + + +@pytest.mark.parametrize( + "base,value", [(str, "x"), (int, 1), (float, 0.5), (bytes, b"x"), (dict, {})] +) +def test_storage_free_scalar_and_mapping_do_not_call_dictionary_property(base, value): + def dictionary(self): + raise AssertionError("no dictionary exists") + + empty = type("Empty", (base,), {"__slots__": (), "__dict__": property(dictionary)}) + options = empty(value) + _tokenize._tokenization_context(options) + with pytest.raises(AttributeError): + options.revision = 0 + + +@pytest.mark.parametrize("base,value", [(str, "x"), (dict, {})]) +@pytest.mark.parametrize("attribute", ["__mro__", "__dict__"]) +@pytest.mark.parametrize("route", ["snapshot", "stable_render", "mutating_render"]) +def test_metaclass_properties_cannot_hide_native_slots(base, value, attribute, route): + accesses = [] + + def hidden(cls): + accesses.append(attribute) + return () if attribute == "__mro__" else {} + + meta = type("Meta", (type,), {attribute: property(hidden)}) + physical = meta("Physical", (base,), {"__slots__": ("revision",)}) + options = physical(value) + options.revision = [0] + if route == "snapshot": + before = _tokenize._tokenization_context([options, options]) + validate = _tokenize._tokenization_context_validator(options) + options.revision[0] = 1 + assert _tokenize._tokenization_context([options, options]) != before + with pytest.raises(ValueError, match="context changed"): + validate(True) + options.revision[0] = 0 + validate(True) + assert _tokenize._tokenization_context([options, options]) == before + else: + _check_callback(options, accesses, route == "mutating_render") + assert accesses == [] From 90ec720d0f969bfc887e322760aedc890e081624 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 10:45:47 +0000 Subject: [PATCH 50/67] Preserve context observer failures and native dictionary evidence --- src/art/trajectories/_tokenize.py | 99 +++++-- .../test_context_observer_errors.py | 251 +++++++++++++++++- .../trajectories/test_scalar_context_state.py | 123 +++++++++ 3 files changed, 454 insertions(+), 19 deletions(-) diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index cf1fda964..8c882c5b3 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -1021,6 +1021,36 @@ def __call__( ) -> None: ... +class _UnsupportedTokenizationContext(TypeError): + """An internal proof refusal, distinct from errors raised by observers.""" + + +class _ContextObserver: + failed: BaseException | None = None + + def __call__( + self, + value: object, + *, + _observed: dict[int, tuple[object, object]] | None = None, + _required: bool = False, + ) -> object: + if self.failed is not None: + raise self.failed + try: + return _tokenization_context(value, _observed=_observed) + except _UnsupportedTokenizationContext as error: + if not _required: + raise + self.failed = ValueError( + "Tokenization context cannot be checked after admission" + ) + raise self.failed from error + except BaseException as error: + self.failed = error + raise + + class _PlainRenderPickler(Pickler): def reducer_override(self, value: object) -> Any: # Exact builtins use the C traversal. Never call custom reducers. @@ -1063,7 +1093,14 @@ def check_value(self, value: object, expected: tuple) -> None: if expected[0] == "plain" else _tokenization_context(value) == expected[1] ) - except (TypeError, ValueError, RecursionError, PicklingError): + except _UnsupportedTokenizationContext: + unchanged = False + except BaseException as error: + if expected[0] != "plain" or not isinstance( + error, (TypeError, ValueError, RecursionError, PicklingError) + ): + self.failed = error + raise unchanged = False if not unchanged: self.failed = ValueError( @@ -1091,7 +1128,13 @@ def reset(self) -> None: def call(self, function: Callable[..., Any], *args: Any, **kwargs: Any) -> Any: self.check() arguments = [list(args), kwargs] - expected = self.snapshot(arguments) + try: + expected = self.snapshot(arguments) + except _UnsupportedTokenizationContext: + raise + except BaseException as error: + self.failed = error + raise try: result = function(*args, **kwargs) except BaseException as error: @@ -4853,7 +4896,9 @@ def snapshot(item: object) -> object: observed[identity] = item, result return result if identity in active: - raise TypeError("Unsupported recursive tokenization context") + raise _UnsupportedTokenizationContext( + "Unsupported recursive tokenization context" + ) active.add(identity) try: return snapshot_compound(item, kind, identity) @@ -4863,7 +4908,9 @@ def snapshot(item: object) -> object: def snapshot_compound(item: object, kind: type, identity: int) -> object: if isinstance(item, Enum): if getattr(item, "__objclass__", kind) is not kind: - raise TypeError("Unsupported enum tokenization context") + raise _UnsupportedTokenizationContext( + "Unsupported enum tokenization context" + ) return kind, instance_state( item, { @@ -4893,6 +4940,15 @@ def snapshot_compound(item: object, kind: type, identity: int) -> object: # A renderer can read mapping attributes as well as its items. # Inspect physical instance storage without invoking overrides. result = (*result, instance_state(item, instance_dictionary(item))) + if isinstance(item, dict): + # Semantic items may hide payload still visible to native lookup. + result = ( + *result, + tuple( + (snapshot(key), snapshot(child)) + for key, child in dict.items(item) + ), + ) elif isinstance(item, Exchange): result = kind, identity, item.model, snapshot(item.request) elif isinstance(item, BaseModel): @@ -4906,7 +4962,9 @@ def snapshot_compound(item: object, kind: type, identity: int) -> object: instance_state(item, instance_dictionary(item)), ) else: - raise TypeError("Unsupported mutable tokenization context") + raise _UnsupportedTokenizationContext( + "Unsupported mutable tokenization context" + ) observed[identity] = item, result return result @@ -4915,9 +4973,16 @@ def instance_dictionary(item: object) -> object: for owner in type.__dict__["__mro__"].__get__(kind): descriptor = type.__dict__["__dict__"].__get__(owner).get("__dict__") if type(descriptor) is GetSetDescriptorType: - return descriptor.__get__(item, kind) + dictionary = descriptor.__get__(item, kind) + if type(dictionary) is not dict: + raise _UnsupportedTokenizationContext( + "Unsupported physical instance dictionary" + ) + return dictionary if type.__dict__["__dictoffset__"].__get__(kind): - raise TypeError("Unsupported hidden instance dictionary") + raise _UnsupportedTokenizationContext( + "Unsupported hidden instance dictionary" + ) return None def instance_state(item: object, dictionary: object) -> object: @@ -4945,9 +5010,10 @@ def instance_state(item: object, dictionary: object) -> object: def _tokenization_context_validator(value: object) -> Callable[[bool], None]: + observe = _ContextObserver() try: - expected = _tokenization_context(value) - except TypeError: + expected = observe(value) + except _UnsupportedTokenizationContext: # An opaque, complete native history still works without a tokenizer. # A callback-bearing path cannot claim an uncheckable context proof. expected = None @@ -4957,7 +5023,7 @@ def validate(require_supported: bool) -> None: return if expected is None: raise ValueError("Tokenization context cannot be checked for callbacks") - if _tokenization_context(value) != expected: + if observe(value, _required=True) != expected: raise ValueError( "Tokenization context changed during tokenization callback" ) @@ -4987,6 +5053,7 @@ def _sampled_source_validator( rendered_evidence: bool = False, _fingerprints: _RevalidatedFingerprints | None = None, ) -> _SampledSourceValidator: + observe = _ContextObserver() expected = {} observed: dict[int, tuple[object, object]] = {} items: Iterable[tuple[_SampledSourceKey, object]] = ( @@ -4999,8 +5066,8 @@ def _sampled_source_validator( if exchange is None: raise ValueError("Sampled source has no exchange") try: - request = _tokenization_context(exchange.request, _observed=observed) - except TypeError: + request = observe(exchange.request, _observed=observed) + except _UnsupportedTokenizationContext: request = None try: plain_request = ( @@ -5018,7 +5085,7 @@ def _sampled_source_validator( _source_stop_evidence(source, key), request, plain_request, - _tokenization_context( + observe( { name: exchange.request.get(name) for name in selected_request_fields @@ -5026,7 +5093,7 @@ def _sampled_source_validator( ) if selected_request_fields is not None else None, - _tokenization_context(_rendered_response_evidence(source)) + observe(_rendered_response_evidence(source)) if rendered_evidence and isinstance(exchange, ResponsesExchange) else None, ) @@ -5071,7 +5138,7 @@ def validate( ) if ( visible is not None - and _tokenization_context(_rendered_response_evidence(source)) + and observe(_rendered_response_evidence(source), _required=True) != visible ): raise ValueError( @@ -5102,7 +5169,7 @@ def validate( } request = selected_request if ( - _tokenization_context(current_request, _observed=observed) + observe(current_request, _observed=observed, _required=True) != request ): raise ValueError( diff --git a/tests/unit/trajectories/test_context_observer_errors.py b/tests/unit/trajectories/test_context_observer_errors.py index d1e2bb95c..f0590c7a0 100644 --- a/tests/unit/trajectories/test_context_observer_errors.py +++ b/tests/unit/trajectories/test_context_observer_errors.py @@ -2,10 +2,10 @@ from pydantic import BaseModel import pytest -from test_tokenize import _chat_exchange +from test_tokenize import _character_template_history, _chat_exchange import art.trajectories as tr -from art.trajectories._tokenize import _tokenization_context +from art.trajectories._tokenize import _RenderContextGuard, _tokenization_context def _observer(kind, failure): @@ -31,7 +31,17 @@ def __getattribute__(self, name): @pytest.mark.parametrize("kind", ["items", "partial_items", "value", "model_extra"]) -@pytest.mark.parametrize("error_type", [RecursionError, RuntimeError]) +@pytest.mark.parametrize( + "error_type", + [ + TypeError, + KeyError, + NotImplementedError, + ValueError, + RecursionError, + RuntimeError, + ], +) @pytest.mark.parametrize("route", ["snapshot", "native", "render"]) def test_observer_failure_is_not_an_unsupported_context(kind, error_type, route): failure = error_type("observer failure") @@ -70,3 +80,238 @@ def invoke(): invoke() assert actual.value is failure assert calls == [] + + +def _one_shot_observer(kind, armed): + def check(): + if armed: + raise armed.pop() + + if kind == "items": + + class Options(dict): + def items(self): + check() + return super().items() + + return Options(key="stable") + + class Model(BaseModel): + value: str = "stable" + + def __getattribute__(self, name): + if name == "value": + check() + return super().__getattribute__(name) + + return Model() + + +@pytest.mark.parametrize("kind", ["items", "model"]) +@pytest.mark.parametrize( + "error_type", + [ + TypeError, + KeyError, + NotImplementedError, + ValueError, + RecursionError, + RuntimeError, + ], +) +@pytest.mark.parametrize("observed", ["context", "arguments"]) +def test_post_callback_observer_failure_is_original_and_sticky( + kind, error_type, observed +): + failure = error_type("one-shot observer failure") + armed = [] + options = _one_shot_observer(kind, armed) + guard = _RenderContextGuard(lambda: options if observed == "context" else []) + calls = [] + + def callback(value): + calls.append("callback") + armed.append(failure) + return value + + with pytest.raises(error_type) as actual: + guard.call(callback, options if observed == "arguments" else None) + assert actual.value is failure + assert armed == [] + # An optional-probe fallback cannot turn a transient observer error into + # permission to accept the now-readable original context. + with pytest.raises(error_type) as again: + guard.call(callback, None) + assert again.value is failure + assert calls == ["callback"] + + +@pytest.mark.parametrize("kind", ["items", "model"]) +@pytest.mark.parametrize( + "error_type", + [ + TypeError, + KeyError, + NotImplementedError, + ValueError, + RecursionError, + RuntimeError, + ], +) +def test_argument_observer_failure_cannot_be_swallowed_by_optional_probe( + kind, error_type +): + failure = error_type("argument observer") + armed = [failure] + options = _one_shot_observer(kind, armed) + guard = _RenderContextGuard(lambda: []) + calls = [] + with pytest.raises(error_type) as actual: + guard.call(lambda value: calls.append(value), options) + assert actual.value is failure and armed == [] + with pytest.raises(error_type) as again: + guard.call(lambda: calls.append("fallback")) + assert again.value is failure and calls == [] + + +@pytest.mark.parametrize("kind", ["items", "model"]) +@pytest.mark.parametrize( + "error_type", + [ + TypeError, + KeyError, + NotImplementedError, + ValueError, + RecursionError, + RuntimeError, + ], +) +@pytest.mark.parametrize("validator_kind", ["history", "sampled"]) +def test_native_validator_cannot_lose_a_one_shot_observer_error( + kind, error_type, validator_kind +): + from art.trajectories import _tokenize + + failure = error_type("native observer") + armed = [] + options = _one_shot_observer(kind, armed) + if validator_kind == "history": + history_check = _tokenize._tokenization_context_validator(options) + + def validate(): + history_check(True) + else: + exchange = _chat_exchange([1], [2]) + exchange.request["metadata"] = cast(Any, {"options": options}) + key = _tokenize._exchange_sampled_source_key(exchange) + source_check = _tokenize._sampled_source_validator({key: exchange}) + + def validate(): + source_check(None) + + armed.append(failure) + with pytest.raises(error_type) as actual: + validate() + assert actual.value is failure and armed == [] + # Native capability handlers revalidate before declining to another path. + with pytest.raises(error_type) as again: + validate() + assert again.value is failure + + +@pytest.mark.parametrize("kind", ["items", "model"]) +@pytest.mark.parametrize( + "error_type", + [ + TypeError, + KeyError, + NotImplementedError, + ValueError, + RecursionError, + RuntimeError, + ], +) +def test_native_boundary_fallback_cannot_hide_request_observer_failure( + monkeypatch, kind, error_type +): + history, tokenizer, _ = _character_template_history() + source = history.message_sources[3] + assert source is not None + armed = [] + failure = error_type("boundary request observer") + source.exchange.request["metadata"] = cast( + Any, {"options": _one_shot_observer(kind, armed)} + ) + decode = tokenizer.decode + calls = [] + + def arm_once(tokens, **kwargs): + calls.append("decode") + if len(calls) == 1: + armed.append(failure) + return decode(tokens, **kwargs) + + monkeypatch.setattr(tokenizer, "decode", arm_once) + with pytest.raises(error_type) as actual: + history.tokenize(tokenizer=tokenizer) + assert actual.value is failure + assert calls == ["decode"] and armed == [] + + +@pytest.mark.parametrize("route", ["history", "sampled", "boundary"]) +def test_transient_unsupported_state_after_admission_cannot_decline(monkeypatch, route): + from art.trajectories import _tokenize + + armed = [] + + class Transient(dict): + __slots__ = () + + def items(self): + if armed: + armed.pop() + return [("temporary", object())] + return [] + + options = Transient() + if route == "boundary": + history, tokenizer, _ = _character_template_history() + source = history.message_sources[3] + assert source is not None + source.exchange.request["metadata"] = cast(Any, {"options": options}) + decode = tokenizer.decode + calls = [] + + def arm_once(tokens, **kwargs): + calls.append("decode") + if len(calls) == 1: + armed.append(True) + return decode(tokens, **kwargs) + + monkeypatch.setattr(tokenizer, "decode", arm_once) + with pytest.raises(ValueError, match="cannot be checked after admission"): + history.tokenize(tokenizer=tokenizer) + assert calls == ["decode"] and armed == [] + return + if route == "history": + history_check = _tokenize._tokenization_context_validator(options) + + def validate(): + history_check(True) + else: + exchange = _chat_exchange([1], [2]) + exchange.request["metadata"] = cast(Any, {"options": options}) + key = _tokenize._exchange_sampled_source_key(exchange) + source_check = _tokenize._sampled_source_validator({key: exchange}) + + def validate(): + source_check(None) + + armed.append(True) + with pytest.raises(ValueError, match="cannot be checked after admission") as first: + validate() + assert armed == [] + assert isinstance(first.value.__cause__, _tokenize._UnsupportedTokenizationContext) + with pytest.raises(ValueError) as again: + validate() + assert again.value is first.value diff --git a/tests/unit/trajectories/test_scalar_context_state.py b/tests/unit/trajectories/test_scalar_context_state.py index bdf8d4ca7..55ee8ae6e 100644 --- a/tests/unit/trajectories/test_scalar_context_state.py +++ b/tests/unit/trajectories/test_scalar_context_state.py @@ -1,6 +1,8 @@ from enum import Enum +from types import GetSetDescriptorType from typing import Any, cast +from pydantic import BaseModel import pytest from test_tokenize import _chat_exchange @@ -63,6 +65,51 @@ def test_scalar_snapshot_tracks_physical_state_and_restoration(kind, override): assert accesses == [] +@pytest.mark.parametrize( + "kind", ["str", "int", "float", "bytes", "enum", "string_enum", "mapping", "model"] +) +@pytest.mark.parametrize("behavior", ["hidden", "raising"]) +def test_rich_physical_dictionary_is_not_a_native_storage_proof(kind, behavior): + reads = [] + + class Dictionary(dict): + def items(self): + reads.append("items") + if behavior == "raising": + raise AssertionError("physical storage must not call rich items") + return [] + + if kind == "mapping": + options = type("Options", (dict,), {})(key="stable") + setattr(options, "revision", [0]) + elif kind == "model": + + class Model(BaseModel): + revision: list[int] = [0] + + options = Model() + else: + options, _ = _scalar(kind, "hidden") + for owner in type.__dict__["__mro__"].__get__(type(options)): + descriptor = type.__dict__["__dict__"].__get__(owner).get("__dict__") + if type(descriptor) is GetSetDescriptorType: + dictionary = Dictionary(descriptor.__get__(options)) + descriptor.__set__(options, dictionary) + break + else: + raise AssertionError("fixture has no native dictionary") + assert getattr(options, "revision") == [0] + getattr(options, "revision")[0] = 1 + assert dict.__getitem__(dictionary, "revision") == [1] + with pytest.raises(TypeError, match="Unsupported physical instance dictionary"): + _tokenize._tokenization_context(options) + validate = _tokenize._tokenization_context_validator(options) + validate(False) + with pytest.raises(ValueError, match="cannot be checked"): + validate(True) + assert reads == [] + + @pytest.mark.parametrize( "kind", ["str", "int", "float", "bytes", "enum", "string_enum"] ) @@ -212,3 +259,79 @@ def hidden(cls): else: _check_callback(options, accesses, route == "mutating_render") assert accesses == [] + + +@pytest.mark.parametrize("view", ["hidden", "alternate", "truthful"]) +@pytest.mark.parametrize("storage", ["dictionary", "slots"]) +@pytest.mark.parametrize("route", ["snapshot", "stable_render", "mutating_render"]) +def test_rich_dict_tracks_native_payload(view, storage, route): + class Options(dict): + __slots__ = () + + @property + def revision(self): + return dict.__getitem__(self, "revision") + + def items(self): + if view == "hidden": + return [] + if view == "alternate": + return [("visible", "stable")] + return dict.items(self) + + kind = ( + type("WithDictionary", (Options,), {}) if storage == "dictionary" else Options + ) + options = kind(revision=[0]) + if route != "snapshot": + _check_callback(options, [], route == "mutating_render") + return + context = [options, options] + before = _tokenize._tokenization_context(context) + validate = _tokenize._tokenization_context_validator(context) + options.revision[0] = 1 + assert _tokenize._tokenization_context(context) != before + with pytest.raises(ValueError, match="context changed"): + validate(True) + options.revision[0] = 0 + assert _tokenize._tokenization_context(context) == before + validate(True) + + +@pytest.mark.parametrize("shape", ["key", "nested", "cycle"]) +def test_rich_dict_cannot_hide_native_key_nested_state_or_cycle(shape): + class Hidden(dict): + __slots__ = () + + def items(self): + return [] + + class Key(str): + revision: list[int] + + key = Key("key") + key.revision = [0] + payload = Hidden({key: [0]}) + if shape == "cycle": + payload["self"] = payload + with pytest.raises(TypeError, match="Unsupported recursive"): + _tokenize._tokenization_context(payload) + return + if shape == "nested": + + class Outer(str): + payload: Hidden + + context = Outer("outer") + context.payload = payload + else: + context = payload + before = _tokenize._tokenization_context(context) + validate = _tokenize._tokenization_context_validator(context) + if shape == "key": + key.revision[0] = 1 + else: + payload[key][0] = 1 + assert _tokenize._tokenization_context(context) != before + with pytest.raises(ValueError, match="context changed"): + validate(True) From 112d30b8e4b9611e9ab6948af5c8f316e585ed7b Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 10:51:57 +0000 Subject: [PATCH 51/67] Use identity for exact type admission and context class tags --- src/art/trajectories/_tokenize.py | 37 +++-- .../trajectories/test_exact_context_types.py | 128 ++++++++++++++++++ 2 files changed, 153 insertions(+), 12 deletions(-) create mode 100644 tests/unit/trajectories/test_exact_context_types.py diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index 8c882c5b3..3521d7b17 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -4891,7 +4891,15 @@ def snapshot(item: object) -> object: previous = observed.get(identity) if previous is not None and previous[0] is item: return previous[1] - if kind in (str, int, bool, bytes, type(None), datetime, float): + if ( + kind is str + or kind is int + or kind is bool + or kind is bytes + or kind is type(None) + or kind is datetime + or kind is float + ): result = kind, repr(item) if kind is float else item observed[identity] = item, result return result @@ -4906,12 +4914,13 @@ def snapshot(item: object) -> object: active.remove(identity) def snapshot_compound(item: object, kind: type, identity: int) -> object: + tag = type_tag(kind) if isinstance(item, Enum): if getattr(item, "__objclass__", kind) is not kind: raise _UnsupportedTokenizationContext( "Unsupported enum tokenization context" ) - return kind, instance_state( + return tag, instance_state( item, { key: child @@ -4928,12 +4937,12 @@ def snapshot_compound(item: object, kind: type, identity: int) -> object: scalar = repr(float.__float__(item)) else: scalar = bytes.__bytes__(item) - return kind, scalar, instance_state(item, instance_dictionary(item)) - if kind in (list, tuple, set, frozenset): - result = kind, tuple(snapshot(child) for child in cast(Iterable, item)) + return tag, scalar, instance_state(item, instance_dictionary(item)) + if kind is list or kind is tuple or kind is set or kind is frozenset: + result = tag, tuple(snapshot(child) for child in cast(Iterable, item)) elif type(item) is dict or isinstance(item, Mapping): result = ( - kind, + tag, tuple((snapshot(key), snapshot(child)) for key, child in item.items()), ) if kind is not dict: @@ -4950,10 +4959,10 @@ def snapshot_compound(item: object, kind: type, identity: int) -> object: ), ) elif isinstance(item, Exchange): - result = kind, identity, item.model, snapshot(item.request) + result = tag, identity, item.model, snapshot(item.request) elif isinstance(item, BaseModel): result = ( - kind, + tag, tuple( (name, snapshot(getattr(item, name))) for name in type(item).model_fields @@ -4968,6 +4977,10 @@ def snapshot_compound(item: object, kind: type, identity: int) -> object: observed[identity] = item, result return result + def type_tag(kind: type) -> object: + # Preserve class lifetime while comparing custom metaclasses by identity. + return kind if type(kind) is type else (id(kind), kind) + def instance_dictionary(item: object) -> object: kind = type(item) for owner in type.__dict__["__mro__"].__get__(kind): @@ -4997,9 +5010,9 @@ def instance_state(item: object, dictionary: object) -> object: try: child = descriptor.__get__(item, kind) except AttributeError: - slots.append((owner, name, False)) + slots.append((type_tag(owner), name, False)) else: - slots.append((owner, name, True, snapshot(child))) + slots.append((type_tag(owner), name, True, snapshot(child))) return snapshot(dictionary), tuple(slots) try: @@ -6286,8 +6299,8 @@ def __init__( ) if ( pure_template - and type(tokenizer_template) in (str, type(None)) - and type(template) in (str, type(None)) + and (type(tokenizer_template) is str or tokenizer_template is None) + and (type(template) is str or template is None) ): # Plain unnamed template normalization has no tokenizer callback. template, normalization_template, defaults = _resolved_chat_template( diff --git a/tests/unit/trajectories/test_exact_context_types.py b/tests/unit/trajectories/test_exact_context_types.py new file mode 100644 index 000000000..a9dde29c2 --- /dev/null +++ b/tests/unit/trajectories/test_exact_context_types.py @@ -0,0 +1,128 @@ +from datetime import datetime +from typing import Any, cast + +import pytest +from test_tokenize import _character_template_history + +from art.trajectories import _tokenize + + +@pytest.mark.parametrize( + "target", + [str, int, bool, bytes, type(None), datetime, float, list, tuple, set, frozenset], +) +@pytest.mark.parametrize("shape", ["opaque", "mapping", "scalar"]) +def test_metaclass_equality_cannot_grant_exact_type_admission(target, shape): + class Meta(type): + def __eq__(cls, other): + return other is target + + __hash__ = type.__hash__ + + class Opaque(metaclass=Meta): + revision: list[int] + + def __iter__(self): + return iter(()) + + class Mapping(dict, metaclass=Meta): + revision: list[int] + + class Scalar(str, metaclass=Meta): + revision: list[int] + + value = {"opaque": Opaque, "mapping": Mapping, "scalar": Scalar}[shape]() + value.revision = [0] + if shape == "opaque": + with pytest.raises(TypeError, match="Unsupported mutable"): + _tokenize._tokenization_context(value) + validate = _tokenize._tokenization_context_validator(value) + validate(False) + with pytest.raises(ValueError, match="cannot be checked"): + validate(True) + else: + expected = _tokenize._tokenization_context(value) + validate = _tokenize._tokenization_context_validator(value) + value.revision[0] = 1 + assert _tokenize._tokenization_context(value) != expected + with pytest.raises(ValueError, match="context changed"): + validate(True) + value.revision[0] = 0 + assert _tokenize._tokenization_context(value) == expected + validate(True) + + +@pytest.mark.parametrize("location", ["tokenizer", "explicit"]) +def test_rich_template_normalization_remains_guarded(monkeypatch, location): + class Meta(type): + def __eq__(cls, other): + return other is str + + __hash__ = type.__hash__ + + class Template(str, metaclass=Meta): + pass + + history, tokenizer, _ = _character_template_history() + calls = [] + original = _tokenize._TraceBuilder.checked + + def checked(self, function, *args, **kwargs): + if function is _tokenize._resolved_chat_template: + calls.append("guarded normalization") + return original(self, function, *args, **kwargs) + + monkeypatch.setattr(_tokenize._TraceBuilder, "checked", checked) + if location == "tokenizer": + setattr(tokenizer, "chat_template", Template("public template")) + result = history.tokenize(tokenizer=tokenizer) + else: + result = history.tokenize( + tokenizer=tokenizer, chat_template=cast(Any, Template("public template")) + ) + assert result.tokens and calls + + +@pytest.mark.parametrize("base", [str, dict]) +def test_metaclass_equality_cannot_hide_renderer_visible_type_replacement(base): + class Meta(type): + def __eq__(cls, other): + return other is Left or other is Right + + __hash__ = type.__hash__ + + Left = Meta("Left", (base,), {}) + Right = Meta("Right", (base,), {}) + left, right = Left(), Right() + left.revision = right.revision = [0] + context = [left] + expected = _tokenize._tokenization_context(context) + validate = _tokenize._tokenization_context_validator(context) + context[0] = right + assert _tokenize._tokenization_context(context) != expected + with pytest.raises(ValueError, match="context changed"): + validate(True) + context[0] = left + assert _tokenize._tokenization_context(context) == expected + validate(True) + + +def test_custom_type_tag_preserves_class_lifetime_without_instance_retention(): + import gc + import weakref + + class Meta(type): + pass + + class Value(str, metaclass=Meta): + pass + + value = Value("public") + kind_ref, value_ref = weakref.ref(Value), weakref.ref(value) + expected = _tokenize._tokenization_context(value) + del value, Value + gc.collect() + assert kind_ref() is not None and value_ref() is None + del expected + gc.collect() + assert kind_ref() is None From f9203d282691c50591e7fc7c196bbd57f186c60e Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 10:54:40 +0000 Subject: [PATCH 52/67] Use identity for plain render-cache type eligibility --- src/art/trajectories/_render_cache.py | 14 ++++-- .../test_render_cache_eligibility.py | 43 +++++++++++++++++++ 2 files changed, 54 insertions(+), 3 deletions(-) diff --git a/src/art/trajectories/_render_cache.py b/src/art/trajectories/_render_cache.py index a78fc7396..92560232c 100644 --- a/src/art/trajectories/_render_cache.py +++ b/src/art/trajectories/_render_cache.py @@ -10,7 +10,7 @@ def _render_context_key(value: object) -> object: """Snapshot plain JSON without losing mapping order or scalar types.""" kind = type(value) - if kind in (str, int, bool, type(None)): + if kind is str or kind is int or kind is bool or kind is type(None): return kind, value if kind is float and math.isfinite(cast(float, value)): return kind, repr(value) @@ -51,7 +51,10 @@ def cacheable_chat_template(tokenizer, template, tools, kwargs, messages) -> boo or base_module.render_jinja_template is not chat.render_jinja_template or not cls.__module__.startswith("transformers.") or getattr(module, cls.__name__, None) is not cls - or type(tokenizer.chat_template) not in (str, type(None)) + or ( + (template_type := type(tokenizer.chat_template)) is not str + and template_type is not type(None) + ) or inspect.getattr_static(cls, "special_tokens_map") is not inspect.getattr_static(base, "special_tokens_map") ): @@ -67,7 +70,12 @@ def cacheable_chat_template(tokenizer, template, tools, kwargs, messages) -> boo added_token = sys.modules["tokenizers"].AddedToken special = tokenizer._special_tokens_map if type(special) is not dict or any( - type(key) is not str or type(value) not in (str, type(None), added_token) + type(key) is not str + or ( + type(value) is not str + and value is not None + and type(value) is not added_token + ) for key, value in special.items() ): return False diff --git a/tests/unit/trajectories/test_render_cache_eligibility.py b/tests/unit/trajectories/test_render_cache_eligibility.py index 56af42446..728cb3076 100644 --- a/tests/unit/trajectories/test_render_cache_eligibility.py +++ b/tests/unit/trajectories/test_render_cache_eligibility.py @@ -199,3 +199,46 @@ def test_generation_extension_override_bypasses(tokenizer, monkeypatch): tracker, "_generation_support", lambda *args, **kwargs: "changed" ) assert not eligible(tokenizer, template) + + +@pytest.mark.parametrize("target", [str, int, bool, type(None)]) +def test_metaclass_equality_cannot_admit_mutable_plain_context(target): + from art.trajectories._render_cache import _render_context_key + + class Meta(type): + def __eq__(cls, other): + return other is target + + __hash__ = type.__hash__ + + class Mutable(metaclass=Meta): + pass + + with pytest.raises(TypeError, match="Not a plain JSON"): + _render_context_key(Mutable()) + + +@pytest.mark.parametrize( + "location", ["chat_template", "special", "messages", "tools", "kwargs"] +) +def test_metaclass_equality_cannot_admit_rich_render_cache_value( + tokenizer, monkeypatch, location +): + class Meta(type): + def __eq__(cls, other): + return other is str + + __hash__ = type.__hash__ + + class Rich(str, metaclass=Meta): + pass + + value = Rich("mutable") + context = {} + if location == "chat_template": + monkeypatch.setattr(tokenizer, "chat_template", value) + elif location == "special": + monkeypatch.setitem(tokenizer._special_tokens_map, "eos_token", value) + else: + context[location] = {"mutable": value} + assert not eligible(tokenizer, "{{ messages }}", **context) From 52ae6fee69caa21e13fc9452052ae74e9b0fc2e6 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 10:56:46 +0000 Subject: [PATCH 53/67] Assert retained identity in custom context type tags --- tests/unit/trajectories/test_context_observation.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/tests/unit/trajectories/test_context_observation.py b/tests/unit/trajectories/test_context_observation.py index 585e6de60..2382cc83c 100644 --- a/tests/unit/trajectories/test_context_observation.py +++ b/tests/unit/trajectories/test_context_observation.py @@ -24,7 +24,11 @@ def test_context_preserves_mapping_order_aliases_and_fresh_mutable_reads(mapping ) first, second = original[1] assert first is second - assert first[0] is mapping_type + if type(mapping_type) is type: + assert first[0] is mapping_type + else: + assert first[0][0] == id(mapping_type) + assert first[0][1] is mapping_type assert [entry[0] for entry in first[1]] == [(str, "first"), (str, "second")] assert first[1][0][1] is first[1][1][1] if isinstance(mapping, ObservedDict): From d09152f406427dd2e0d11d66642977bf6008c0dc Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 11:13:04 +0000 Subject: [PATCH 54/67] Release sticky observer failures when tokenization ends --- src/art/trajectories/_tokenize.py | 76 ++++++-- .../test_failed_observer_lifetime.py | 179 ++++++++++++++++++ 2 files changed, 242 insertions(+), 13 deletions(-) create mode 100644 tests/unit/trajectories/test_failed_observer_lifetime.py diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index 3521d7b17..4ddb645b4 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -3,11 +3,12 @@ from bisect import bisect_left import codecs from collections.abc import Callable, Iterable, Mapping, Sequence +from contextvars import ContextVar from copy import deepcopy from dataclasses import dataclass, field, replace from datetime import datetime from enum import Enum -from functools import lru_cache +from functools import lru_cache, wraps from hashlib import sha256 from inspect import getattr_static from io import BytesIO @@ -1025,6 +1026,49 @@ class _UnsupportedTokenizationContext(TypeError): """An internal proof refusal, distinct from errors raised by observers.""" +@dataclass +class _FailureScope: + active: bool = True + states: list[_ContextObserver | _RenderContextGuard] = field(default_factory=list) + + +_FAILURE_SCOPE: ContextVar[_FailureScope | None] = ContextVar( + "art_tokenization_failure_scope", default=None +) + + +def _record_context_failure( + state: _ContextObserver | _RenderContextGuard, error: BaseException +) -> BaseException: + state.failed = error + scope = _FAILURE_SCOPE.get() + if scope is not None and scope.active: + scope.states.append(state) + return error + + +def _release_context_failures[**P, R](function: Callable[P, R]) -> Callable[P, R]: + @wraps(function) + def owned(*args: P.args, **kwargs: P.kwargs) -> R: + outer = _FAILURE_SCOPE.get() + if outer is not None and outer.active: + return function(*args, **kwargs) + scope = _FailureScope() + token = _FAILURE_SCOPE.set(scope) + try: + return function(*args, **kwargs) + finally: + # Nested histories must retain sticky failures through final validation. + # Restore ownership before releasing objects whose finalizers may reenter. + scope.active = False + _FAILURE_SCOPE.reset(token) + for state in scope.states: + state.failed = None + scope.states.clear() + + return owned + + class _ContextObserver: failed: BaseException | None = None @@ -1042,12 +1086,12 @@ def __call__( except _UnsupportedTokenizationContext as error: if not _required: raise - self.failed = ValueError( - "Tokenization context cannot be checked after admission" - ) - raise self.failed from error + raise _record_context_failure( + self, + ValueError("Tokenization context cannot be checked after admission"), + ) from error except BaseException as error: - self.failed = error + _record_context_failure(self, error) raise @@ -1099,14 +1143,14 @@ def check_value(self, value: object, expected: tuple) -> None: if expected[0] != "plain" or not isinstance( error, (TypeError, ValueError, RecursionError, PicklingError) ): - self.failed = error + _record_context_failure(self, error) raise unchanged = False if not unchanged: - self.failed = ValueError( - "Rendering context changed during tokenization callback" + raise _record_context_failure( + self, + ValueError("Rendering context changed during tokenization callback"), ) - raise self.failed def check(self) -> None: if self.failed is not None: @@ -1114,7 +1158,7 @@ def check(self) -> None: try: value = self.read() except BaseException as error: - self.failed = error + _record_context_failure(self, error) raise self.check_value(value, self.expected) @@ -1133,7 +1177,7 @@ def call(self, function: Callable[..., Any], *args: Any, **kwargs: Any) -> Any: except _UnsupportedTokenizationContext: raise except BaseException as error: - self.failed = error + _record_context_failure(self, error) raise try: result = function(*args, **kwargs) @@ -1144,7 +1188,7 @@ def call(self, function: Callable[..., Any], *args: Any, **kwargs: Any) -> Any: except BaseException: # Optional renderer probes can catch the original exception. # Keep its identity and prevent a later fallback blessing edits. - self.failed = error + _record_context_failure(self, error) raise self.check() self.check_value(arguments, expected) @@ -5022,6 +5066,7 @@ def instance_state(item: object, dictionary: object) -> object: del snapshot, snapshot_compound, instance_state +@_release_context_failures def _tokenization_context_validator(value: object) -> Callable[[bool], None]: observe = _ContextObserver() try: @@ -5058,6 +5103,7 @@ def _rendered_response_evidence(source: object) -> object: ) +@_release_context_failures def _sampled_source_validator( sources: Mapping[_SampledSourceKey, object] | Sequence[tuple[_SampledSourceKey, object]], @@ -9138,6 +9184,7 @@ def _tokenize_history( ) +@_release_context_failures def tokenize_history( history: History | LegacyHistory, *, @@ -9304,6 +9351,7 @@ def _complete_resolved_sampled_stops( _validate_completed_sources(builders) +@_release_context_failures def tokenize_trajectory( trajectory: Trajectory, *, @@ -9373,6 +9421,7 @@ def tokenize_trajectory( ) +@_release_context_failures def _tokenize_trajectory_with_trace( trajectory: Trajectory, *, @@ -9428,6 +9477,7 @@ def _tokenize_trajectory_with_trace( ) +@_release_context_failures def tokenize_group( group: TrajectoryGroup, *, diff --git a/tests/unit/trajectories/test_failed_observer_lifetime.py b/tests/unit/trajectories/test_failed_observer_lifetime.py new file mode 100644 index 000000000..45f7dfff4 --- /dev/null +++ b/tests/unit/trajectories/test_failed_observer_lifetime.py @@ -0,0 +1,179 @@ +import gc +from typing import Any, cast +import weakref + +import pytest +from test_tokenize import _character_template_history, _chat_exchange + +import art.trajectories as tr +from art.trajectories import _tokenize + + +@pytest.mark.parametrize( + "route", ["native", "boundary", "render", "history_factory", "sampled_factory"] +) +@pytest.mark.parametrize("error_type", [TypeError, RuntimeError]) +def test_failed_tokenization_releases_observer_inputs(route, error_type): + enabled = gc.isenabled() + gc.disable() + try: + + def fail(): + armed = ( + [True] + if route in ("native", "history_factory", "sampled_factory") + else [] + ) + + class Options(dict): + def items(self): + if armed: + armed.pop() + raise error_type("public observer failure") + return dict.items(self) + + options = Options(key="stable") + reference = weakref.ref(options) + try: + if route == "history_factory": + _tokenize._tokenization_context_validator(options) + elif route == "sampled_factory": + exchange = _chat_exchange([1], [2]) + exchange.request["metadata"] = cast(Any, {"options": options}) + key = _tokenize._exchange_sampled_source_key(exchange) + _tokenize._sampled_source_validator({key: exchange}) + elif route == "boundary": + history, tokenizer, _ = _character_template_history() + source = history.message_sources[3] + assert source is not None + source.exchange.request["metadata"] = cast( + Any, {"options": options} + ) + decode = tokenizer.decode + + def arm(tokens, **kwargs): + armed.append(True) + return decode(tokens, **kwargs) + + setattr(tokenizer, "decode", arm) + history.tokenize(tokenizer=tokenizer) + else: + exchange = _chat_exchange([1], [2]) + value = tr.Trajectory( + exchanges=tr.TrajectoryExchanges(chat_completions=[exchange]) + ) + if route == "native": + exchange.request["metadata"] = cast(Any, {"options": options}) + value.tokenize() + else: + + class Tokenizer: + def apply_chat_template(self, messages, **kwargs): + armed.append(True) + return [1, 2] + + def __call__(self, text, **kwargs): + return [2 if text == "answer" else 1] + + value.tokenize( + tokenizer=cast(Any, Tokenizer()), + chat_template="custom", + chat_template_kwargs={"options": options}, + ) + except error_type: + return reference + pytest.fail("observer unexpectedly accepted") + + reference = fail() + assert reference() is None + finally: + if enabled: + gc.enable() + gc.collect() + + +def test_explicit_private_owner_keeps_failure_sticky_until_operation_ends(): + enabled = gc.isenabled() + gc.disable() + try: + + def operation(): + armed = [] + + class Options(dict): + def items(self): + if armed: + armed.pop() + raise TypeError("one-shot") + return dict.items(self) + + options = Options(key="stable") + reference = weakref.ref(options) + validate = _tokenize._tokenization_context_validator(options) + armed.append(True) + try: + validate(True) + except TypeError as first: + with pytest.raises(TypeError) as second: + validate(True) + assert second.value is first + del second + else: + pytest.fail("Observer did not fail") + return reference + + reference = _tokenize._release_context_failures(operation)() + assert reference() is None + finally: + if enabled: + gc.enable() + gc.collect() + + +def test_cleanup_finalizer_reentry_has_an_independent_failure_owner(): + enabled = gc.isenabled() + gc.disable() + references = [] + scopes = [] + try: + + def inner(): + class Options(dict): + def items(self): + raise TypeError("reentrant observer") + + options = Options() + references.append(weakref.ref(options)) + scopes.append(_tokenize._FAILURE_SCOPE.get()) + _tokenize._tokenization_context_validator(options) + + class ReenteringError(TypeError): + def __del__(self): + try: + _tokenize._release_context_failures(inner)() + except TypeError: + pass + + def outer(): + class Options(dict): + def items(self): + raise ReenteringError("outer observer") + + scopes.append(_tokenize._FAILURE_SCOPE.get()) + try: + _tokenize._tokenization_context_validator(Options()) + except TypeError: + pass + + _tokenize._release_context_failures(outer)() + assert len(scopes) == 2 and scopes[0] is not scopes[1] + assert all( + scope is not None and not scope.active and not scope.states + for scope in scopes + ) + assert len(references) == 1 and references[0]() is None + assert _tokenize._FAILURE_SCOPE.get() is None + finally: + if enabled: + gc.enable() + gc.collect() From 74f0554eb05540487c97f85006ab007849f4b5fa Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 11:19:12 +0000 Subject: [PATCH 55/67] Preserve shared trim bindings across calls and parser fallback --- src/art_inference/chat_template.py | 18 ++++++- tests/unit/test_literal_reasoning_content.py | 50 ++++++++++++++++++++ 2 files changed, 67 insertions(+), 1 deletion(-) diff --git a/src/art_inference/chat_template.py b/src/art_inference/chat_template.py index 3c85c0e06..7c6d58f21 100644 --- a/src/art_inference/chat_template.py +++ b/src/art_inference/chat_template.py @@ -137,7 +137,19 @@ def apply_edits() -> str: except TemplateSyntaxError: # Optional binding analysis must not undo independently proved parser # edits when the renderer supports extensions absent from this parser. - return apply_edits() + result = apply_edits() + if ( + "" in template + and "reasoning_content|trim" in template + and "if not preserve_thinking or message.reasoning_content" not in template + ): + # Keep the pre-existing preservation rewrite when custom syntax + # prevents the narrower binding analysis below. + result = result.replace( + "set content = render_content(message.content, true)|trim", + "set content = (render_content(message.content, true) if preserve_thinking and message.role == 'assistant' else render_content(message.content, true)|trim)", + ) + return result assignments = list(tree.find_all(nodes.Assign)) locations = [] parsed_assignments = [] @@ -191,6 +203,10 @@ def writes_content(node: nodes.Assign | nodes.AssignBlock) -> bool: ) def reads_content(node: nodes.Node) -> bool: + # A macro/context-aware callable may capture this binding without a + # Name at the call site. Do not move its trim across an unknown call. + if isinstance(node, nodes.Call) or any(node.find_all(nodes.Call)): + return True return ( isinstance(node, nodes.Name) and node.name == "content" diff --git a/tests/unit/test_literal_reasoning_content.py b/tests/unit/test_literal_reasoning_content.py index 26e5afb40..61a1ebb47 100644 --- a/tests/unit/test_literal_reasoning_content.py +++ b/tests/unit/test_literal_reasoning_content.py @@ -242,6 +242,56 @@ def test_unconfigured_template_receives_the_same_correction(): assert chat_template_with_preserved_thinking(raw) == _FIXED +@pytest.mark.parametrize("indirect", [False, True]) +@pytest.mark.parametrize("content", [" answer ", " beforexafter "]) +def test_macro_capture_keeps_shared_content_trim(indirect, content): + parser = _QWEN_INLINE_REASONING.search(_TEMPLATE) + assert parser is not None + template = ( + "{% macro render_content(content, count) %}{{ content }}{% endmacro %}" + "{% macro preview() %}[{{ content }}]{% endmacro %}" + "{% macro wrapper() %}{{ preview() }}{% endmacro %}" + "{% set content = render_content(message.content, true)|trim %}" + + ("{{ wrapper() }}" if indirect else "{{ preview() }}") + + parser.group() + + "[{{ content }}]" + ) + fixed = chat_template_with_preserved_thinking(template) + env = ImmutableSandboxedEnvironment(trim_blocks=True, lstrip_blocks=True) + assert ( + env.from_string(fixed).render(message={"role": "assistant", "content": content}) + == "[" + content.strip() + "][" + content.strip() + "]" + ) + assert chat_template_with_preserved_thinking(fixed) == fixed + + +@pytest.mark.parametrize("preserve", [False, True]) +def test_extension_fallback_keeps_prior_structured_content_whitespace(preserve): + template = "{% for item in [1] %}{% break %}{% endfor %}" + _TEMPLATE.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 %}", + ) + fixed = chat_template_with_preserved_thinking(template) + rendered = _render( + fixed, + [ + _USER, + { + "role": "assistant", + "content": " answer ", + "reasoning_content": "reasoned\n", + }, + ], + enable_thinking=False, + preserve_thinking=preserve, + ) + assert rendered.endswith((" answer " if preserve else "answer") + "<|im_end|>\n") + assert chat_template_with_preserved_thinking(fixed) == 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 From 886f23300dc3c11b1db80e8fef16b6b565fcdb5d Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 11:19:43 +0000 Subject: [PATCH 56/67] Narrow the rendered-template type in the macro regression --- tests/unit/test_literal_reasoning_content.py | 1 + 1 file changed, 1 insertion(+) diff --git a/tests/unit/test_literal_reasoning_content.py b/tests/unit/test_literal_reasoning_content.py index 61a1ebb47..947baadb5 100644 --- a/tests/unit/test_literal_reasoning_content.py +++ b/tests/unit/test_literal_reasoning_content.py @@ -257,6 +257,7 @@ def test_macro_capture_keeps_shared_content_trim(indirect, content): + "[{{ content }}]" ) fixed = chat_template_with_preserved_thinking(template) + assert isinstance(fixed, str) env = ImmutableSandboxedEnvironment(trim_blocks=True, lstrip_blocks=True) assert ( env.from_string(fixed).render(message={"role": "assistant", "content": content}) From 759613c64c6f56582c6d395937e83afcf5157e83 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 11:25:55 +0000 Subject: [PATCH 57/67] Require actual model types for context source authority --- src/art/trajectories/_tokenize.py | 4 +- .../test_context_model_admission.py | 86 +++++++++++++++++++ 2 files changed, 88 insertions(+), 2 deletions(-) create mode 100644 tests/unit/trajectories/test_context_model_admission.py diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index 4ddb645b4..4298da509 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -5002,9 +5002,9 @@ def snapshot_compound(item: object, kind: type, identity: int) -> object: for key, child in dict.items(item) ), ) - elif isinstance(item, Exchange): + elif isinstance(item, Exchange) and issubclass(kind, Exchange): result = tag, identity, item.model, snapshot(item.request) - elif isinstance(item, BaseModel): + elif isinstance(item, BaseModel) and issubclass(kind, BaseModel): result = ( tag, tuple( diff --git a/tests/unit/trajectories/test_context_model_admission.py b/tests/unit/trajectories/test_context_model_admission.py new file mode 100644 index 000000000..f563b66f6 --- /dev/null +++ b/tests/unit/trajectories/test_context_model_admission.py @@ -0,0 +1,86 @@ +from typing import Any, cast + +from pydantic import BaseModel +import pytest + +from art.trajectories import ( + ChatCompletionsExchange, + CompletionsExchange, + MessagesExchange, + ResponsesExchange, + _tokenize, +) + +_EXCHANGES = ( + ChatCompletionsExchange, + CompletionsExchange, + MessagesExchange, + ResponsesExchange, +) + + +@pytest.mark.parametrize("target", [*_EXCHANGES, BaseModel]) +@pytest.mark.parametrize("nested", [False, True]) +def test_spoofed_model_class_cannot_admit_unobserved_state(target, nested): + class Fake: + __pydantic_decorators__ = True + # Even a complete model-shaped facade has no actual model authority. + model_fields = {} + model_extra = None + + def __init__(self): + self.model = "public-model" + self.request = {"messages": []} + self.revision = [0] + + @property + def __class__(self): + return target + + value = Fake() + assert isinstance(value, target) + assert not issubclass(type(value), target) + context = [value, value] if nested else value + with pytest.raises(TypeError, match="Unsupported mutable"): + _tokenize._tokenization_context(context) + validate = _tokenize._tokenization_context_validator(context) + validate(False) + value.revision[0] = 1 + with pytest.raises(ValueError, match="cannot be checked"): + validate(True) + + +@pytest.mark.parametrize("target", _EXCHANGES) +@pytest.mark.parametrize("subclass", [False, True]) +def test_actual_exchange_classes_preserve_request_authority(target, subclass): + kind = type("ActualExchange", (target,), {}) if subclass else target + value = cast( + Any, kind.model_construct(request={"model": "public-model", "revision": [0]}) + ) + context = [value, value] + expected = _tokenize._tokenization_context(context) + validate = _tokenize._tokenization_context_validator(context) + value.request["revision"][0] = 1 + assert _tokenize._tokenization_context(context) != expected + with pytest.raises(ValueError, match="context changed"): + validate(True) + value.request["revision"][0] = 0 + assert _tokenize._tokenization_context(context) == expected + validate(True) + + +def test_actual_model_retains_field_and_physical_state_validation(): + class Model(BaseModel): + revision: list[int] + + value = Model(revision=[0]) + context = [value, value] + expected = _tokenize._tokenization_context(context) + validate = _tokenize._tokenization_context_validator(context) + value.revision[0] = 1 + assert _tokenize._tokenization_context(context) != expected + with pytest.raises(ValueError, match="context changed"): + validate(True) + value.revision[0] = 0 + assert _tokenize._tokenization_context(context) == expected + validate(True) From 8bfb7d63d9f7e1a47a74d535c7ca8ab51fbb7d3e Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 11:40:02 +0000 Subject: [PATCH 58/67] Conservatively preserve shared content across unknown template effects --- src/art_inference/chat_template.py | 27 ++-- tests/unit/test_literal_reasoning_content.py | 120 ++++++++++++++++++ .../test_reasoning_parser_joined_branches.py | 4 +- 3 files changed, 135 insertions(+), 16 deletions(-) diff --git a/src/art_inference/chat_template.py b/src/art_inference/chat_template.py index 7c6d58f21..5d145d48d 100644 --- a/src/art_inference/chat_template.py +++ b/src/art_inference/chat_template.py @@ -203,22 +203,21 @@ def writes_content(node: nodes.Assign | nodes.AssignBlock) -> bool: ) def reads_content(node: nodes.Node) -> bool: - # A macro/context-aware callable may capture this binding without a - # Name at the call site. Do not move its trim across an unknown call. - if isinstance(node, nodes.Call) or any(node.find_all(nodes.Call)): - return True - return ( - isinstance(node, nodes.Name) - and node.name == "content" - and node.ctx == "load" - ) or any( - n.name == "content" and n.ctx == "load" for n in node.find_all(nodes.Name) + # Only a plain constant assignment is binding-transparent. Calls, + # filters/tests, loaders, attributes/items, operators and even output + # conversion/finalization may invoke code that observes this scope. + # Keep the trim for unknown nodes instead of enumerating callbacks. + return not ( + isinstance(node, nodes.Assign) + and isinstance(node.target, nodes.Name) + and node.target.name != "message" + and isinstance(node.node, nodes.Const) ) def assistant_condition(test: nodes.Node) -> bool | None: # The rewrite changes assistant content only. Ignore paths proved to - # handle a different role, but inspect every unknown branch for users - # of the original trimmed value before the destructive parser. + # handle a different role in an ordinary chat message dictionary. + # Every unknown condition retains the original trimmed binding. if ( isinstance(test, nodes.Compare) and test.expr @@ -252,9 +251,9 @@ def visit(body: Sequence[nodes.Node], bindings: set[int]) -> set[int]: # If does not introduce a Jinja scope. Retain every binding # reaching the join, including paths that skipped the parser. for branch in (node, *node.elif_): - if reads_content(branch.test): - shared.update(bindings) condition = assistant_condition(branch.test) + if condition is None: + shared.update(bindings) if condition is not False: joined.update(visit(branch.body, bindings)) if condition is True: diff --git a/tests/unit/test_literal_reasoning_content.py b/tests/unit/test_literal_reasoning_content.py index 947baadb5..19ed160e8 100644 --- a/tests/unit/test_literal_reasoning_content.py +++ b/tests/unit/test_literal_reasoning_content.py @@ -2,6 +2,7 @@ import hashlib from pathlib import Path +from jinja2 import DictLoader, pass_context from jinja2.sandbox import ImmutableSandboxedEnvironment import pytest @@ -34,6 +35,125 @@ ) +@pytest.mark.parametrize( + "middle", + [ + "{{ ''|peek }}", + "{% if '' is peek %}seen{% endif %}", + "{% include 'preview' %}", + "{% import 'module' as p with context %}{{ p.stored }}", + "{% from 'module' import stored with context %}{{ stored }}", + "{% set observed = probe.value %}", + "{% set observed = probe['value'] %}", + "{{ probe }}", + "{% if probe %}seen{% endif %}", + "{% set observed = probe + 1 %}", + "{% if probe == 1 %}seen{% endif %}", + "{% for item in probe %}seen{% endfor %}", + "{{ 42 }}", + ], +) +@pytest.mark.parametrize("content", [" answer ", " beforexafter "]) +def test_implicit_context_consumers_keep_original_shared_trim(middle, content): + match = _QWEN_INLINE_REASONING.search(_TEMPLATE) + assert match is not None + parser = match.group() + trim = "{% set content = render_content(message.content, true)|trim %}" + template = ( + "{% macro render_content(content, count) %}{{ content }}{% endmacro %}" + "{% set probe = bind() %}" + trim + middle + parser + "[{{ content }}]" + ) + seen = [] + + class Probe: + def __init__(self, context): + self.context = context + + def observe(self): + value = self.context["content"] + seen.append(value) + return value + + @property + def value(self): + return self.observe() + + def __getitem__(self, key): + return self.observe() + + def __str__(self): + return self.observe() + + def __bool__(self): + self.observe() + return True + + def __add__(self, other): + return self.observe() + + def __eq__(self, other): + self.observe() + return True + + def __iter__(self): + self.observe() + return iter([1]) + + @pass_context + def peek(context, value): + seen.append(context["content"]) + return context["content"] + + @pass_context + def finalize(context, value): + if value == 42: + seen.append(context["content"]) + return value + + env = ImmutableSandboxedEnvironment( + loader=DictLoader( + {"preview": "{{ content }}", "module": "{% set stored=content %}"} + ), + finalize=finalize, + ) + bind = pass_context(lambda context: Probe(context)) + env.filters["peek"] = env.tests["peek"] = peek + fixed = chat_template_with_preserved_thinking(template) + assert isinstance(fixed, str) + assert not _QWEN_INLINE_REASONING.search(fixed) + message = {"role": "assistant", "content": content} + # Removing the destructive parser is intentional; all surrounding consumers + # must continue observing the same original trimmed binding. + expected = env.from_string(template.replace(parser, "")).render( + message=message, bind=bind + ) + before = seen[:] + seen.clear() + assert env.from_string(fixed).render(message=message, bind=bind) == expected + assert seen == before + assert all(value == content.strip() for value in seen) + assert trim in fixed + + +def test_role_guard_is_not_proof_after_message_reassignment(): + match = _QWEN_INLINE_REASONING.search(_TEMPLATE) + assert match is not None + trim = "{% set content = render_content(message.content, true)|trim %}" + template = ( + "{% macro render_content(content, count) %}{{ content }}{% endmacro %}" + + trim + + "{% set message = none %}{% if message.role == 'assistant' %}" + + match.group() + + "{% else %}[{{ content }}]{% endif %}" + ) + fixed = chat_template_with_preserved_thinking(template) + assert isinstance(fixed, str) + env = ImmutableSandboxedEnvironment() + message = {"role": "assistant", "content": " answer "} + assert env.from_string(fixed).render(message=message) == "[answer]" + assert trim in fixed + + @pytest.mark.parametrize("trim_blocks,lstrip_blocks", [(False, False), (True, True)]) @pytest.mark.parametrize( "left,right", [("", ""), ("-", ""), ("", "-"), ("-", "-"), ("+", "+")] diff --git a/tests/unit/test_reasoning_parser_joined_branches.py b/tests/unit/test_reasoning_parser_joined_branches.py index 06a6b0f8b..4c0aa0a01 100644 --- a/tests/unit/test_reasoning_parser_joined_branches.py +++ b/tests/unit/test_reasoning_parser_joined_branches.py @@ -62,7 +62,7 @@ def test_joined_preview_retains_its_original_trim(scope, layout, content): @pytest.mark.parametrize("mode", ["a", "b", "c"]) -def test_all_joined_paths_consuming_parser_can_preserve_assistant_whitespace(mode): +def test_unknown_comparison_keeps_shared_trim_even_when_all_paths_have_parser(mode): match = _QWEN_INLINE_REASONING.search(_TEMPLATE) assert match is not None parser = match.group() @@ -84,6 +84,6 @@ def test_all_joined_paths_consuming_parser_can_preserve_assistant_whitespace(mod env.from_string(fixed).render( mode=mode, message={"role": "assistant", "content": content} ) - == "[" + content + "]" + == "[" + content.strip() + "]" ) assert _without_inline_reasoning_parser(fixed) == fixed From 6d3682f66c9e83ecd7f6eb882402aee7c1afc586 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 11:55:10 +0000 Subject: [PATCH 59/67] Limit transparent assignments to fresh loop-local bindings --- src/art_inference/chat_template.py | 56 +++++++++++++++++--- tests/unit/test_literal_reasoning_content.py | 37 +++++++++++++ 2 files changed, 85 insertions(+), 8 deletions(-) diff --git a/src/art_inference/chat_template.py b/src/art_inference/chat_template.py index 5d145d48d..2c807829e 100644 --- a/src/art_inference/chat_template.py +++ b/src/art_inference/chat_template.py @@ -202,14 +202,17 @@ def writes_content(node: nodes.Assign | nodes.AssignBlock) -> bool: or any(n.name == "content" for n in target.find_all(nodes.Name)) ) - def reads_content(node: nodes.Node) -> bool: - # Only a plain constant assignment is binding-transparent. Calls, + def reads_content(node: nodes.Node, initialized: set[str] | None) -> bool: + # Only a fresh loop-local constant assignment is transparent. Rebinding + # an arbitrary old value can invoke its destructor. Calls, # filters/tests, loaders, attributes/items, operators and even output # conversion/finalization may invoke code that observes this scope. # Keep the trim for unknown nodes instead of enumerating callbacks. return not ( - isinstance(node, nodes.Assign) + initialized is not None + and isinstance(node, nodes.Assign) and isinstance(node.target, nodes.Name) + and node.target.name not in initialized and node.target.name != "message" and isinstance(node.node, nodes.Const) ) @@ -230,7 +233,11 @@ def assistant_condition(test: nodes.Node) -> bool | None: return equal if test.ops[0].op == "eq" else not equal return None - def visit(body: Sequence[nodes.Node], bindings: set[int]) -> set[int]: + def visit( + body: Sequence[nodes.Node], + bindings: set[int], + initialized: set[str] | None = None, + ) -> set[int]: bindings = bindings.copy() for node in body: if isinstance(node, nodes.If): @@ -246,6 +253,12 @@ def visit(body: Sequence[nodes.Node], bindings: set[int]) -> set[int]: else: shared.update(bindings) # No unique consumed assignment. bindings.clear() + if initialized is not None: + initialized.update( + n.name + for n in node.find_all(nodes.Name) + if n.ctx == "store" + ) continue joined = set() # If does not introduce a Jinja scope. Retain every binding @@ -255,15 +268,28 @@ def visit(body: Sequence[nodes.Node], bindings: set[int]) -> set[int]: if condition is None: shared.update(bindings) if condition is not False: - joined.update(visit(branch.body, bindings)) + joined.update(visit(branch.body, bindings, initialized)) if condition is True: break else: - joined.update(visit(node.else_, bindings)) + joined.update(visit(node.else_, bindings, initialized)) bindings = joined else: - if reads_content(node): + if reads_content(node, initialized): shared.update(bindings) + if initialized is not None: + initialized.update( + n.name for n in node.find_all(nodes.Name) if n.ctx == "store" + ) + if isinstance(node, nodes.Macro): + initialized.add(node.name) + elif isinstance(node, nodes.Import): + initialized.add(node.target) + elif isinstance(node, nodes.FromImport): + initialized.update( + name if isinstance(name, str) else name[1] + for name in node.names + ) if isinstance(node, nodes.Assign): if writes_content(node): bindings = {id(node)} @@ -273,7 +299,21 @@ def visit(body: Sequence[nodes.Node], bindings: set[int]) -> set[int]: if isinstance(value, list) and all( isinstance(n, nodes.Node) for n in value ): - visit(value, set()) + # Jinja initializes loop locals before each body; + # parameters/targets already have arbitrary values. + local_names = ( + { + n.name + for n in ( + node.target, + *node.target.find_all(nodes.Name), + ) + if isinstance(n, nodes.Name) + } + if isinstance(node, nodes.For) and value is node.body + else None + ) + visit(value, set(), local_names) if isinstance(node, nodes.AssignBlock) and writes_content(node): bindings.clear() return bindings diff --git a/tests/unit/test_literal_reasoning_content.py b/tests/unit/test_literal_reasoning_content.py index 19ed160e8..c2819e065 100644 --- a/tests/unit/test_literal_reasoning_content.py +++ b/tests/unit/test_literal_reasoning_content.py @@ -51,6 +51,7 @@ "{% if probe == 1 %}seen{% endif %}", "{% for item in probe %}seen{% endfor %}", "{{ 42 }}", + "{% set probe = none %}", ], ) @pytest.mark.parametrize("content", [" answer ", " beforexafter "]) @@ -99,6 +100,10 @@ def __iter__(self): self.observe() return iter([1]) + def __del__(self): + if middle == "{% set probe = none %}": + self.observe() + @pass_context def peek(context, value): seen.append(context["content"]) @@ -154,6 +159,38 @@ def test_role_guard_is_not_proof_after_message_reassignment(): assert trim in fixed +@pytest.mark.parametrize("prior_binding", [False, True]) +def test_only_fresh_loop_local_initialization_is_transparent(prior_binding): + match = _QWEN_INLINE_REASONING.search(_TEMPLATE) + assert match is not None + trim = "{% set content = render_content(message.content, true)|trim %}" + template = ( + "{% macro render_content(content, count) %}{{ content }}{% endmacro %}" + "{% for message in messages %}" + + ("{% set probe = make_probe() %}" if prior_binding else "") + + trim + + "{% set probe = none %}" + + match.group() + + "[{{ content }}]{% endfor %}" + ) + destroyed = [] + + class Probe: + def __del__(self): + destroyed.append(True) + + fixed = chat_template_with_preserved_thinking(template) + assert isinstance(fixed, str) + content = " beforexafter " + env = ImmutableSandboxedEnvironment() + actual = env.from_string(fixed).render( + messages=[{"role": "assistant", "content": content}], make_probe=Probe + ) + assert actual == "[" + (content.strip() if prior_binding else content) + "]" + assert bool(destroyed) is prior_binding + assert (trim in fixed) is prior_binding + + @pytest.mark.parametrize("trim_blocks,lstrip_blocks", [(False, False), (True, True)]) @pytest.mark.parametrize( "left,right", [("", ""), ("-", ""), ("", "-"), ("-", "-"), ("+", "+")] From c5f0d6aee0e928285f80232934cc0b42a4e97bba Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 12:14:35 +0000 Subject: [PATCH 60/67] Track prior stores independently of template role pruning --- src/art_inference/chat_template.py | 43 +++++++++-------- tests/unit/test_literal_reasoning_content.py | 49 ++++++++++++++++++++ 2 files changed, 73 insertions(+), 19 deletions(-) diff --git a/src/art_inference/chat_template.py b/src/art_inference/chat_template.py index 2c807829e..499ee1799 100644 --- a/src/art_inference/chat_template.py +++ b/src/art_inference/chat_template.py @@ -233,6 +233,25 @@ def assistant_condition(test: nodes.Node) -> bool | None: return equal if test.ops[0].op == "eq" else not equal return None + def remember_stores(node: nodes.Node, initialized: set[str] | None) -> None: + if initialized is None: + return + initialized.update( + n.name for n in node.find_all(nodes.Name) if n.ctx == "store" + ) + for child in ( + node, + *node.find_all((nodes.Macro, nodes.Import, nodes.FromImport)), + ): + if isinstance(child, nodes.Macro): + initialized.add(child.name) + elif isinstance(child, nodes.Import): + initialized.add(child.target) + elif isinstance(child, nodes.FromImport): + initialized.update( + name if isinstance(name, str) else name[1] for name in child.names + ) + def visit( body: Sequence[nodes.Node], bindings: set[int], @@ -253,12 +272,7 @@ def visit( else: shared.update(bindings) # No unique consumed assignment. bindings.clear() - if initialized is not None: - initialized.update( - n.name - for n in node.find_all(nodes.Name) - if n.ctx == "store" - ) + remember_stores(node, initialized) continue joined = set() # If does not introduce a Jinja scope. Retain every binding @@ -274,22 +288,13 @@ def visit( else: joined.update(visit(node.else_, bindings, initialized)) bindings = joined + # A role-pruned path may still have initialized a local before + # rebinding message. Freshness follows every syntactic store. + remember_stores(node, initialized) else: if reads_content(node, initialized): shared.update(bindings) - if initialized is not None: - initialized.update( - n.name for n in node.find_all(nodes.Name) if n.ctx == "store" - ) - if isinstance(node, nodes.Macro): - initialized.add(node.name) - elif isinstance(node, nodes.Import): - initialized.add(node.target) - elif isinstance(node, nodes.FromImport): - initialized.update( - name if isinstance(name, str) else name[1] - for name in node.names - ) + remember_stores(node, initialized) if isinstance(node, nodes.Assign): if writes_content(node): bindings = {id(node)} diff --git a/tests/unit/test_literal_reasoning_content.py b/tests/unit/test_literal_reasoning_content.py index c2819e065..b1f3fa9a9 100644 --- a/tests/unit/test_literal_reasoning_content.py +++ b/tests/unit/test_literal_reasoning_content.py @@ -191,6 +191,55 @@ def __del__(self): assert (trim in fixed) is prior_binding +@pytest.mark.parametrize( + "branch", + [ + "{% if message.role == 'user' %}BODY{% endif %}", + "{% if message.role == 'assistant' %}{% else %}BODY{% endif %}", + "{% if message.role == 'system' %}{% elif message.role == 'user' %}BODY{% endif %}", + ], +) +def test_role_pruning_keeps_prior_store_history(branch): + match = _QWEN_INLINE_REASONING.search(_TEMPLATE) + assert match is not None + seen = [] + + class Probe: + def __init__(self, reader): + self.reader = reader + + def __del__(self): + seen.append(self.reader()) + + trim = "{% set content = render_content(message.content, true)|trim %}" + template = ( + "{% macro render_content(content, count) %}{{ content }}{% endmacro %}" + "{% for message in messages %}" + "{% macro read_content() %}{{ content }}{% endmacro %}" + + branch.replace( + "BODY", + "{% set probe = make_probe(read_content) %}" + "{% set message = {'role': 'assistant', 'content': message.content} %}", + ) + + trim + + "{% set probe = none %}" + + match.group() + + "{% endfor %}{{ seen|join('|') }}" + ) + env = ImmutableSandboxedEnvironment() + kwargs = dict( + messages=[{"role": "user", "content": " answer "}], + make_probe=Probe, + seen=seen, + ) + assert env.from_string(template).render(**kwargs) == "answer" + seen.clear() + fixed = _without_inline_reasoning_parser(template) + assert env.from_string(fixed).render(**kwargs) == "answer" + assert seen == ["answer"] + assert trim in fixed + + @pytest.mark.parametrize("trim_blocks,lstrip_blocks", [(False, False), (True, True)]) @pytest.mark.parametrize( "left,right", [("", ""), ("-", ""), ("", "-"), ("-", "-"), ("+", "+")] From 3e8962f866aac69ec1b8f6ad11a5c4a67cf8c070 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 12:40:31 +0000 Subject: [PATCH 61/67] Validate original exchange evidence between template callbacks --- src/art/trajectories/_tokenize.py | 21 +++- ...est_exchange_template_source_boundaries.py | 96 +++++++++++++++++++ 2 files changed, 116 insertions(+), 1 deletion(-) create mode 100644 tests/unit/trajectories/test_exchange_template_source_boundaries.py diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index 4298da509..d6be772c7 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -1195,6 +1195,13 @@ def call(self, function: Callable[..., Any], *args: Any, **kwargs: Any) -> Any: return result +class _SourceRenderGuard(_RenderContextGuard): + """Validate original evidence while allowing disposable renderer arguments.""" + + def call(self, function: Callable[..., Any], *args: Any, **kwargs: Any) -> Any: + return super().call(lambda: function(*args, **kwargs)) + + def _plain_tokenizer_attribute(tokenizer: object, name: str) -> tuple[bool, object]: kind = type(tokenizer) if type(kind) is not type: @@ -3186,7 +3193,19 @@ def checked(function: Callable[..., Any], *args: Any, **kwargs: Any) -> Any: # Selection and rendering are separate callbacks. The retained # projection must remain stable before either consumes it. args = ( - _RenderingTokenizer(args[0], _RenderContextGuard(lambda: projection)), + _RenderingTokenizer( + args[0], + _RenderContextGuard(lambda: ledger.checked(lambda: projection)), + ), + *args[1:], + ) + elif function is _template_ids: + # Conversion makes disposable messages, but tool schemas can still + # alias original requests. Check between selection and rendering. + args = ( + _RenderingTokenizer( + args[0], _SourceRenderGuard(lambda: ledger.checked(lambda: None)) + ), *args[1:], ) if not callback_used: diff --git a/tests/unit/trajectories/test_exchange_template_source_boundaries.py b/tests/unit/trajectories/test_exchange_template_source_boundaries.py new file mode 100644 index 000000000..fa88be833 --- /dev/null +++ b/tests/unit/trajectories/test_exchange_template_source_boundaries.py @@ -0,0 +1,96 @@ +from typing import Any, cast + +import pytest +from test_tokenize import _message_exchange, _response_with_content_logprobs + +import art.trajectories as tr + + +@pytest.mark.parametrize("mutate", [False, True]) +@pytest.mark.parametrize("mutate_copy", [False, True]) +def test_template_selection_cannot_change_original_schema_then_restore( + mutate, mutate_copy +): + exchange = _message_exchange( + cast( + Any, + { + "model": "test/model", + "max_tokens": 5, + "messages": [{"role": "user", "content": "question"}], + "tools": [{"name": "ask", "input_schema": {"type": "string"}}], + }, + ), + token_ids=[2], + logprobs=[-0.2], + ) + events = [] + + class Tokenizer: + chat_template = {"named": "unchanged body"} + + def get_chat_template(self, chat_template=None, tools: Any = None): + schema = tools[0]["function"]["parameters"] + assert schema is cast(Any, exchange.request)["tools"][0]["input_schema"] + events.append("select") + if mutate: + schema["type"] = "integer" + return "unchanged body" + + def apply_chat_template(self, messages, tools: Any = None, **kwargs): + events.append("render") + schema = tools[0]["function"]["parameters"] + changed = schema["type"] == "integer" + schema["type"] = "string" + if mutate_copy: + assert messages is not exchange.request["messages"] + messages[0]["content"] = "disposable change" + return [99 if changed else 1] + + trajectory = tr.Trajectory(exchanges=tr.TrajectoryExchanges(messages=[exchange])) + if mutate: + with pytest.raises(ValueError, match="[Cc]ontext changed"): + trajectory.tokenize(tokenizer=cast(Any, Tokenizer())) + assert events == ["select"] + else: + result = trajectory.tokenize(tokenizer=cast(Any, Tokenizer())) + assert result.tokens == [1, 2] + assert result.logprobs[-1] == -0.2 + assert result.flags[-1] & tr.TokenFlag.SAMPLED + assert exchange.request["messages"][0]["content"] == "question" + assert events == ["select", "render"] + + +@pytest.mark.parametrize("mutate", [False, True]) +def test_responses_template_selection_validates_original_source(mutate): + exchange = _response_with_content_logprobs(exact_second=True) + original = exchange.request["input"] + events = [] + + class Tokenizer: + chat_template = {"named": "unchanged body"} + + def get_chat_template(self, chat_template=None, tools=None): + events.append("select") + if mutate: + exchange.request["input"] = "changed original" + return "unchanged body" + + def apply_chat_template(self, messages, tools=None, **kwargs): + events.append("render") + changed = exchange.request["input"] != original + exchange.request["input"] = original + return [99 if changed else 1] + + trajectory = tr.Trajectory(exchanges=tr.TrajectoryExchanges(responses=[exchange])) + if mutate: + with pytest.raises(ValueError, match="[Cc]ontext changed"): + trajectory.tokenize(tokenizer=cast(Any, Tokenizer())) + assert events == ["select"] + else: + result = trajectory.tokenize(tokenizer=cast(Any, Tokenizer())) + assert result.tokens == [1, 11, 12] + assert result.logprobs[1:] == [-0.1, -0.2] + assert all(flag & tr.TokenFlag.SAMPLED for flag in result.flags[1:]) + assert exchange.request["input"] == original + assert events == ["select", "render"] From cca25f4159a4ab57a2e6e314ab517fb2dbcf8c76 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 12:44:09 +0000 Subject: [PATCH 62/67] Refuse unavailable historical request role proofs consistently --- docs/features/additional-histories.mdx | 9 ++++- src/art/trajectories/_tokenize.py | 2 + .../trajectories/test_recorded_boundaries.py | 24 ++++++++++++ .../test_recorded_prompt_roles.py | 39 +++++++++++++++++++ 4 files changed, 72 insertions(+), 2 deletions(-) diff --git a/docs/features/additional-histories.mdx b/docs/features/additional-histories.mdx index 76dc2597a..4081dc726 100644 --- a/docs/features/additional-histories.mdx +++ b/docs/features/additional-histories.mdx @@ -181,7 +181,9 @@ For supported Chat boundaries, ART decodes the recorded body and encodes only th 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. +continue through the generic rendering path and its source validation. That +fallback can still refuse copied response context when it cannot prove the +original conditioning of the sampled tokens that follow it. Correcting literal-content rendering does not rewrite a recorded request. When the original request and template reproduce its complete native prompt, ART can @@ -247,7 +249,10 @@ 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 +that coverage; tokenize the containing trajectory with `multi_history=True` +and keep `reconcile_text_equivalent_tokenizations=False` (the default) so the +complete original occurrence remains in a selected history. Reconciliation can +collapse that owner into a shortened view, which ART safely refuses. This correction does not change how generic output/SFT masks include copied assistant content. ## How It Works diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index d6be772c7..2c7880a9c 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -6909,6 +6909,7 @@ def _substitute_exact_prefix(self) -> None: signature = _source_signature(source) request_context = None full_prompt_proven = False + request_messages = None try: # Canonical tool validation may reorder JSON keys. # Prove the historical prompt with the recorded @@ -6947,6 +6948,7 @@ def _substitute_exact_prefix(self) -> None: if ( needs_request_roles and self.recorded_prompt_masks is None + and request_messages is not None ): try: raw_prompt = _recorded_prompt_tokens( diff --git a/tests/unit/trajectories/test_recorded_boundaries.py b/tests/unit/trajectories/test_recorded_boundaries.py index 0e4af63c9..d475b3427 100644 --- a/tests/unit/trajectories/test_recorded_boundaries.py +++ b/tests/unit/trajectories/test_recorded_boundaries.py @@ -923,3 +923,27 @@ def render(*args, **kwargs): _trace=None, ) assert not called + + +def test_reconciliation_cannot_remove_complete_copied_sample_owner() -> 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) + second.request["messages"] = [ + {"role": "user", "content": "turn 0"}, + {"role": "assistant", "content": "answer"}, + {"role": "user", "content": "turn 1"}, + ] + trajectory = tr.Trajectory( + exchanges=tr.TrajectoryExchanges(chat_completions=[first, second]) + ) + result = trajectory.tokenize(multi_history=True) + assert [item.tokens for item in result.histories] == [[1, 2, 3], [1, 3, 4, 5]] + assert len(trajectory.histories(reconcile_text_equivalent_tokenizations=True)) == 1 + with pytest.raises(ValueError, match="complete original sampled occurrence"): + trajectory.tokenize( + multi_history=True, reconcile_text_equivalent_tokenizations=True + ) diff --git a/tests/unit/trajectories/test_recorded_prompt_roles.py b/tests/unit/trajectories/test_recorded_prompt_roles.py index d425f042b..42cc42f1c 100644 --- a/tests/unit/trajectories/test_recorded_prompt_roles.py +++ b/tests/unit/trajectories/test_recorded_prompt_roles.py @@ -3,6 +3,7 @@ from copy import deepcopy import json import math +from types import SimpleNamespace from typing import Any, cast from openai.types.chat import ChatCompletionMessageParam @@ -630,3 +631,41 @@ def test_named_original_request_proof_normalizes_tool_arguments(selection): flag & tr.TokenFlag.SAMPLED for flag in history.flags[len(prompt) : len(prompt) + len(output)] ) + + +@pytest.mark.parametrize("failure", [TypeError, KeyError, NotImplementedError]) +def test_unavailable_original_request_refuses_role_proof_without_unbound_locals( + monkeypatch: pytest.MonkeyPatch, failure: type[Exception] +) -> None: + trajectory, _, _, _ = _case("Historical assistant") + history = trajectory.chat_completions_history() + calls = [] + + def unavailable(exchange): + calls.append(exchange) + raise failure("Unavailable historical request projection") + + monkeypatch.setattr(module, "_request_messages", unavailable) + stage = SimpleNamespace( + rendered=[1, 2], + chat_template=None, + chat_template_kwargs=None, + projection_matches=True, + history=history, + messages=history.messages, + canonical_assistant_mask=[True], + canonical_stop_mask=[False], + original_template="template", + template="template", + validate_consumed=lambda _: None, + _source_prompt_tokens=lambda source: ( + [9] if module._source_is_sampled(source) else None + ), + _source_matches_context=lambda _: True, + _probe_render=lambda *args, **kwargs: [1], + prompt_cache={}, + output_cache={}, + ) + with pytest.raises(ValueError, match="Cannot preserve request roles"): + module._ChatViewTokenizer._substitute_exact_prefix(cast(Any, stage)) + assert len(calls) == 1 From 3d3726934f55e69b83c342f35e9b0423dd90c9e0 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 12:57:45 +0000 Subject: [PATCH 63/67] Keep shared trim for unproved renderer bindings and context effects --- src/art_inference/chat_template.py | 180 ++++++++++++- tests/unit/test_literal_reasoning_content.py | 263 ++++++++++++++++++- 2 files changed, 437 insertions(+), 6 deletions(-) diff --git a/src/art_inference/chat_template.py b/src/art_inference/chat_template.py index 499ee1799..8ed44f769 100644 --- a/src/art_inference/chat_template.py +++ b/src/art_inference/chat_template.py @@ -46,7 +46,7 @@ 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, nodes + from jinja2 import Environment, TemplateSyntaxError, meta, nodes from jinja2.visitor import NodeTransformer # Compare parsed operations, not quote/spacing choices or a template hash. @@ -61,6 +61,9 @@ def visit_Output(self, node: nodes.Output, *args: Any, **kwargs: Any): return node env = Environment() + # This environment only analyzes source. Treat even built-in globals as + # external bindings rather than hiding them from Jinja's scope analysis. + env.globals.clear() def operations(text: str): return WithoutWhitespace().visit(env.parse(text)).body @@ -150,6 +153,147 @@ def apply_edits() -> str: "set content = (render_content(message.content, true) if preserve_thinking and message.role == 'assistant' else render_content(message.content, true)|trim)", ) return result + + # Moving the role test ahead of a renderer is safe only for a local, + # data-only macro under the stock chat renderer. External callbacks and + # unknown macro effects retain the original trim and evaluation order. + # This permits ordinary chat data, built-in tests, integer vision counters + # and the stock raising helper, not custom finalizers or renderer globals. + renderers = [ + node for node in tree.find_all(nodes.Macro) if node.name == "render_content" + ] + if len(renderers) != 1 or renderers[0] not in tree.body: + return apply_edits() + if any(tree.find_all(nodes.Extends)): + return apply_edits() # A parent template can replace the local macro. + renderer = renderers[0] + preceding = tree.body[: tree.body.index(renderer)] + if any( + name.name == "render_content" + for node in preceding + for name in node.find_all(nodes.Name) + ): + return apply_edits() + isolated = nodes.Template([renderer]) + isolated.set_environment(env) + try: + # Use Jinja's own scope analysis, including conditionally bound locals. + global_names = meta.find_undeclared_variables(isolated) + except TemplateSyntaxError: + return apply_edits() + counters = global_names - {"raise_exception", "add_vision_id"} + for name in counters: + declarations = [ + node + for node in preceding + if isinstance(node, nodes.Assign) + and node.target == nodes.Name(name, "store") + ] + declaration = env.parse("{% set counter = namespace(value=0) %}").body[0] + assert isinstance(declaration, nodes.Assign) + expected = declaration.node + if len(declarations) != 1 or declarations[0].node != expected: + return apply_edits() + inside = {id(node) for node in renderer.find_all(nodes.Node)} + if any( + node.name == name + and id(node) not in inside + and (isinstance(node, nodes.NSRef) or node.ctx == "load") + for node in tree.find_all((nodes.Name, nodes.NSRef)) + ): + return apply_edits() + allowed = tuple( + getattr(nodes, name) + for name in ( + "Add And Assign Call Compare Concat Const For Getattr If Name Not " + "NSRef Operand Or Output TemplateData Test" + ).split() + ) + protected = { + "render_content", + "raise_exception", + "namespace", + "add_vision_id", + *counters, + } + if any( + not isinstance(node, allowed) + or isinstance(node, nodes.Call) + and not ( + node.node == nodes.Name("raise_exception", "load") + and len(node.args) == 1 + and isinstance(node.args[0], nodes.Const) + and not node.kwargs + and node.dyn_args is None + and node.dyn_kwargs is None + ) + or isinstance(node, nodes.Test) + and node.name not in {"string", "iterable", "mapping", "none", "undefined"} + for node in renderer.find_all(nodes.Node) + ): + return apply_edits() + if any( + node.name in protected - counters + and node.ctx != "load" + or node.name in counters + and node.ctx != "load" + and not any( + node is declaration.target + for declaration in tree.body + if isinstance(declaration, nodes.Assign) + ) + for node in tree.find_all(nodes.Name) + ) or any( + isinstance(node, nodes.Import) + and node.target in protected + or isinstance(node, nodes.FromImport) + and any( + (name if isinstance(name, str) else name[1]) in protected + for name in node.names + ) + or isinstance(node, nodes.Macro) + and node is not renderer + and node.name in protected + for node in tree.find_all((nodes.Import, nodes.FromImport, nodes.Macro)) + ): + return apply_edits() + if counters: + # Namespace counters are mutable. A context callback elsewhere could + # replace their values without a lexical reference to the namespace. + # Admit only the stock chat template's closed data/render operations. + closed = WithoutWhitespace().visit(env.parse(apply_edits())) + try: + external = meta.find_undeclared_variables(closed) + except TemplateSyntaxError: + return apply_edits() + if external.difference( + "messages tools message content reasoning_content preserve_thinking " + "enable_thinking add_generation_prompt add_vision_id raise_exception namespace".split() + ): + return apply_edits() + data_nodes = allowed + tuple( + getattr(nodes, name) + for name in "CondExpr Filter Getitem Keyword Macro Neg Slice Sub Tuple".split() + ) + for node in closed.find_all(nodes.Node): + if not isinstance(node, data_nodes): + return apply_edits() + if isinstance(node, nodes.Call) and not ( + isinstance(node.node, nodes.Name) + and node.node.name in {"render_content", "raise_exception", "namespace"} + or isinstance(node.node, nodes.Getattr) + and node.node.node == nodes.Name("content", "load") + and node.node.attr in {"startswith", "endswith"} + ): + return apply_edits() + if isinstance(node, nodes.Filter) and node.name not in ( + "default items length safe string tojson trim".split() + ): + return apply_edits() + if isinstance(node, nodes.Test) and node.name not in ( + "defined false iterable mapping none string true undefined".split() + ): + return apply_edits() assignments = list(tree.find_all(nodes.Assign)) locations = [] parsed_assignments = [] @@ -256,8 +400,13 @@ def visit( body: Sequence[nodes.Node], bindings: set[int], initialized: set[str] | None = None, + loop_locals: bool = False, ) -> set[int]: bindings = bindings.copy() + # Store history belongs to every scope; only loop locals have the + # ownership proof allowing a fresh constant assignment to be ignored. + if initialized is None: + initialized = set() for node in body: if isinstance(node, nodes.If): if (id(node) in edited_parsers and node == operation[0]) or ( @@ -282,22 +431,29 @@ def visit( if condition is None: shared.update(bindings) if condition is not False: - joined.update(visit(branch.body, bindings, initialized)) + joined.update( + visit(branch.body, bindings, initialized, loop_locals) + ) if condition is True: break else: - joined.update(visit(node.else_, bindings, initialized)) + joined.update(visit(node.else_, bindings, initialized, loop_locals)) bindings = joined # A role-pruned path may still have initialized a local before # rebinding message. Freshness follows every syntactic store. remember_stores(node, initialized) else: - if reads_content(node, initialized): + if reads_content(node, initialized if loop_locals else None): shared.update(bindings) + replacing_content = bool(bindings) or ("content" in initialized) remember_stores(node, initialized) if isinstance(node, nodes.Assign): if writes_content(node): bindings = {id(node)} + if replacing_content: + # Publishing a new binding can release an old + # object whose destructor observes the new value. + shared.update(bindings) else: # Macro/loop/with/block bodies have independent bindings. for _, value in node.iter_fields(): @@ -318,7 +474,21 @@ def visit( if isinstance(node, nodes.For) and value is node.body else None ) - visit(value, set(), local_names) + if isinstance(node, nodes.Macro): + local_names = {arg.name for arg in node.args} + elif isinstance(node, nodes.With): + local_names = { + bound.name + for target in node.targets + for bound in (target, *target.find_all(nodes.Name)) + if isinstance(bound, nodes.Name) + } + visit( + value, + set(), + local_names, + isinstance(node, nodes.For) and value is node.body, + ) if isinstance(node, nodes.AssignBlock) and writes_content(node): bindings.clear() return bindings diff --git a/tests/unit/test_literal_reasoning_content.py b/tests/unit/test_literal_reasoning_content.py index b1f3fa9a9..f9d1ec2e0 100644 --- a/tests/unit/test_literal_reasoning_content.py +++ b/tests/unit/test_literal_reasoning_content.py @@ -240,6 +240,223 @@ def __del__(self): assert trim in fixed +@pytest.mark.parametrize( + "prefix", + [ + "", + "{% macro render_content(content, count) %}{{ mutate(content, count) }}{% endmacro %}", + "{% macro render_content(content, count) %}{{ content|mutate }}{% endmacro %}", + "{% macro render_content(content, count) %}{% if content is mutate %}{{ content }}{% endif %}{% endmacro %}", + "{% macro render_content(content, count) %}{{ probe.value }}{% endmacro %}", + "{% set dict = probe %}{% macro render_content(content, count) %}{{ dict.value }}{% endmacro %}", + "{% macro render_content(content, count) %}{{ probe.value }}{% if false %}{% set probe = 0 %}{% endif %}{% endmacro %}", + "{% macro render_content(content, count) %}{{ probe['value'] }}{% endmacro %}", + "{% macro render_content(content, count) %}{{ content }}{% endmacro %}{% set render_content = mutate %}", + ], +) +@pytest.mark.parametrize("mutate_role", [False, True]) +def test_unknown_renderer_effects_keep_trim_and_call_order(prefix, mutate_role): + match = _QWEN_INLINE_REASONING.search(_TEMPLATE) + assert match is not None + trim = "{% set content = render_content(message.content, true)|trim %}" + source = ( + prefix + + trim + + "{% if message.role == 'assistant' %}" + + match.group() + + "{% endif %}[{{ content }}]" + ) + env = ImmutableSandboxedEnvironment() + + def render(template): + message = {"role": "assistant", "content": " beforexafter "} + calls = [] + + def mutate(value, count=None): + calls.append(message["role"]) + if mutate_role: + message["role"] = "user" + return value + + class Probe: + @property + def value(self): + return mutate(message["content"]) + + def __getitem__(self, key): + return self.value + + env.filters["mutate"] = mutate + env.tests["mutate"] = lambda value: bool(mutate(value)) + result = env.from_string(template).render( + message=message, render_content=mutate, mutate=mutate, probe=Probe() + ) + return result, calls + + fixed = _without_inline_reasoning_parser(source) + assert not _QWEN_INLINE_REASONING.search(fixed) + assert trim in fixed + assert render(fixed) == render(source.replace(match.group(), "")) + assert render(fixed) == ("[beforexafter]", ["assistant"]) + assert chat_template_with_preserved_thinking(source) == fixed + assert chat_template_with_preserved_thinking(fixed) == fixed + + +def test_renderer_declared_after_use_does_not_prove_role_stability(): + match = _QWEN_INLINE_REASONING.search(_TEMPLATE) + assert match is not None + trim = "{% set content = render_content(message.content, true)|trim %}" + source = ( + trim + + "{% if message.role == 'assistant' %}" + + match.group() + + "{% else %}[{{ content }}]{% endif %}" + "{% macro render_content(content, count) %}{{ content }}{% endmacro %}" + ) + env = ImmutableSandboxedEnvironment() + + def render(template): + message = {"role": "assistant", "content": " answer "} + + def render_content(value, count): + message["role"] = "user" + return value + + return env.from_string(template).render( + message=message, render_content=render_content + ) + + fixed = _without_inline_reasoning_parser(source) + assert trim in fixed + assert render(source) == render(fixed) == "[answer]" + + +def test_inherited_renderer_does_not_prove_role_stability(): + match = _QWEN_INLINE_REASONING.search(_TEMPLATE) + assert match is not None + trim = "{% set content = render_content(message.content, true)|trim %}" + source = ( + "{% extends 'parent' %}" + "{% macro render_content(content,count) %}{{ content }}{% endmacro %}" + "{% block body %}" + + trim + + "{% if message.role == 'assistant' %}" + + match.group() + + "{% endif %}[{{ content }}]{% endblock %}" + ) + parent = ( + "{% macro render_content(content,count) %}{{ mutate(content,count) }}{% endmacro %}" + "{% block body %}{% endblock %}" + ) + message = {"role": "assistant", "content": " answer "} + + def mutate(value, count): + message["role"] = "user" + return value + + fixed = _without_inline_reasoning_parser(source) + assert trim in fixed + assert ( + ImmutableSandboxedEnvironment(loader=DictLoader({"parent": parent})) + .from_string(fixed) + .render(message=message, mutate=mutate) + == "[answer]" + ) + + +@pytest.mark.parametrize( + "middle", + [ + "{% include 'setup' %}", + "{{ inject() }}", + "{{ ''|inject }}", + "{% if '' is inject %}{% endif %}", + ], +) +def test_counter_renderer_rejects_implicit_context_exports(middle): + match = _QWEN_INLINE_REASONING.search(_TEMPLATE) + assert match is not None + trim = "{% set content = render_content(message.content, true)|trim %}" + source = ( + "{% set counter=namespace(value=0) %}" + "{% macro render_content(content,count) %}{% set counter.value=counter.value+1 %}{{ content }}{% endmacro %}" + + middle + + trim + + "{% if message.role == 'assistant' %}" + + match.group() + + "{% endif %}[{{ content }}]" + ) + message = {"role": "assistant", "content": " answer "} + + class Probe: + def __add__(self, other): + message["role"] = "user" + return 1 + + @pass_context + def inject(context, value=None): + context["counter"]["value"] = Probe() + return "" + + env = ImmutableSandboxedEnvironment( + loader=DictLoader({"setup": "{% set counter.value=probe %}"}) + ) + env.filters["inject"] = env.tests["inject"] = inject + fixed = _without_inline_reasoning_parser(source) + assert trim in fixed + assert ( + env.from_string(fixed).render(message=message, inject=inject, probe=Probe()) + == "[answer]" + ) + + +@pytest.mark.parametrize("scope", ["top", "loop", "macro", "with"]) +def test_replacing_content_keeps_trim_for_its_own_observers(scope): + match = _QWEN_INLINE_REASONING.search(_TEMPLATE) + assert match is not None + seen = [] + + class Probe: + def __init__(self, reader): + self.reader = reader + + def __del__(self): + seen.append(self.reader()) + + trim = "{% set content = render_content(message.content,true)|trim %}" + body = ( + "{% macro read_content() %}{{ content }}{% endmacro %}" + "{% if message.role == 'user' %}{% set content = make_probe(read_content) %}" + "{% set message = {'role': 'assistant', 'content': message.content} %}{% endif %}" + + trim + + match.group() + ) + if scope == "loop": + body = "{% for message in messages %}" + body + "{% endfor %}" + elif scope == "macro": + body = "{% macro run(message) %}" + body + "{% endmacro %}{{ run(message) }}" + elif scope == "with": + body = "{% with message=message %}" + body + "{% endwith %}" + source = ( + "{% macro render_content(content,count) %}{{ content }}{% endmacro %}" + + body + + "{{ seen|join('|') }}" + ) + fixed = _without_inline_reasoning_parser(source) + assert trim in fixed + assert ( + ImmutableSandboxedEnvironment() + .from_string(fixed) + .render( + messages=[{"role": "user", "content": " answer "}], + message={"role": "user", "content": " answer "}, + make_probe=Probe, + seen=seen, + ) + == "answer" + ) + + @pytest.mark.parametrize("trim_blocks,lstrip_blocks", [(False, False), (True, True)]) @pytest.mark.parametrize( "left,right", [("", ""), ("-", ""), ("", "-"), ("-", "-"), ("+", "+")] @@ -758,7 +975,9 @@ def test_qwen_content_preview_macro_is_unchanged(preserve, raw): [_USER, {"role": "assistant", "content": content}], preserve_thinking=preserve, ) - assert content + "<|im_end|>\n" in rendered + # The extra renderer reference before its declaration prevents the closed + # stock-template proof. Keep its existing preservation/trim behavior. + assert (content if preserve else content.strip()) + "<|im_end|>\n" in rendered assert rendered.endswith("[beforeliteralafter]") @@ -829,6 +1048,48 @@ def parse(self, parser): assert chat_template_with_preserved_thinking(fixed) == fixed +def test_tuple_scope_content_replacement_retains_trim(): + match = _QWEN_INLINE_REASONING.search(_TEMPLATE) + assert match is not None + seen = [] + + class Probe: + def __init__(self): + self.reader = lambda: "unbound" + + def __del__(self): + seen.append(self.reader()) + + def bind(probe, reader): + probe.reader = reader + return "" + + trim = "{% set content = render_content(message.content,true)|trim %}" + source = ( + "{% macro render_content(content,count) %}{{ content }}{% endmacro %}" + "{% with (content, spare)=make_pair() %}" + "{% macro read_content() %}{{ content }}{% endmacro %}" + "{{ bind(content,read_content) }}" + + trim + + match.group() + + "{% endwith %}{{ seen|join('|') }}" + ) + fixed = _without_inline_reasoning_parser(source) + assert trim in fixed + assert ( + ImmutableSandboxedEnvironment() + .from_string(fixed) + .render( + message={"role": "assistant", "content": " answer "}, + make_pair=lambda: (Probe(), None), + bind=bind, + seen=seen, + ) + == "answer" + ) + assert seen == ["answer"] + + @pytest.mark.parametrize("consumer", ["output", "alias", "condition", "other_branch"]) def test_shared_content_preview_prevents_ambiguous_trim_rewrite(consumer): match = _QWEN_INLINE_REASONING.search(_TEMPLATE) From 4aed0d43185866029f361d7e80fa57c5b65a27ee Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 13:22:04 +0000 Subject: [PATCH 64/67] Expect the earlier consumed-projection refusal in response controls --- tests/unit/trajectories/test_reused_response_selection.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/tests/unit/trajectories/test_reused_response_selection.py b/tests/unit/trajectories/test_reused_response_selection.py index 1cdb0cdca..0ae6392fb 100644 --- a/tests/unit/trajectories/test_reused_response_selection.py +++ b/tests/unit/trajectories/test_reused_response_selection.py @@ -138,7 +138,9 @@ def test_responses_selection_does_not_change_reused_projection( assert rendered[0]["consumed"] == "turn 0" if mode in ("consumed_then_restored", "lasting"): assert row["error"]["class"] == "ValueError" - assert "context changed" in row["error"]["message"] + assert row["error"]["message"] == ( + "Consumed source text changed during tokenization callback" + ) # The changed projection never reaches the completion renderer. assert len(rendered) == 1 else: From f9d834e2483c54e1733a950021054c5548bdf55e Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 13:19:15 +0000 Subject: [PATCH 65/67] Check whole-template effects for counter-free renderers --- src/art_inference/chat_template.py | 67 +++++++++--------- tests/unit/test_literal_reasoning_content.py | 71 ++++++++++++++++++-- 2 files changed, 98 insertions(+), 40 deletions(-) diff --git a/src/art_inference/chat_template.py b/src/art_inference/chat_template.py index 8ed44f769..5e1daed6d 100644 --- a/src/art_inference/chat_template.py +++ b/src/art_inference/chat_template.py @@ -257,43 +257,42 @@ def apply_edits() -> str: for node in tree.find_all((nodes.Import, nodes.FromImport, nodes.Macro)) ): return apply_edits() - if counters: - # Namespace counters are mutable. A context callback elsewhere could - # replace their values without a lexical reference to the namespace. - # Admit only the stock chat template's closed data/render operations. - closed = WithoutWhitespace().visit(env.parse(apply_edits())) - try: - external = meta.find_undeclared_variables(closed) - except TemplateSyntaxError: + # Macro callables and namespace counters are mutable. A context callback + # can replace either without a lexical reference, even for counter-free macros. + # Admit only the stock chat template's closed data/render operations. + closed = WithoutWhitespace().visit(env.parse(apply_edits())) + try: + external = meta.find_undeclared_variables(closed) + except TemplateSyntaxError: + return apply_edits() + if external.difference( + "messages tools message content reasoning_content preserve_thinking " + "enable_thinking add_generation_prompt add_vision_id raise_exception namespace".split() + ): + return apply_edits() + data_nodes = allowed + tuple( + getattr(nodes, name) + for name in "CondExpr Filter Getitem Keyword Macro Neg Slice Sub Tuple".split() + ) + for node in closed.find_all(nodes.Node): + if not isinstance(node, data_nodes): return apply_edits() - if external.difference( - "messages tools message content reasoning_content preserve_thinking " - "enable_thinking add_generation_prompt add_vision_id raise_exception namespace".split() + if isinstance(node, nodes.Call) and not ( + isinstance(node.node, nodes.Name) + and node.node.name in {"render_content", "raise_exception", "namespace"} + or isinstance(node.node, nodes.Getattr) + and node.node.node == nodes.Name("content", "load") + and node.node.attr in {"startswith", "endswith"} + ): + return apply_edits() + if isinstance(node, nodes.Filter) and node.name not in ( + "default items length safe string tojson trim".split() + ): + return apply_edits() + if isinstance(node, nodes.Test) and node.name not in ( + "defined false iterable mapping none string true undefined".split() ): return apply_edits() - data_nodes = allowed + tuple( - getattr(nodes, name) - for name in "CondExpr Filter Getitem Keyword Macro Neg Slice Sub Tuple".split() - ) - for node in closed.find_all(nodes.Node): - if not isinstance(node, data_nodes): - return apply_edits() - if isinstance(node, nodes.Call) and not ( - isinstance(node.node, nodes.Name) - and node.node.name in {"render_content", "raise_exception", "namespace"} - or isinstance(node.node, nodes.Getattr) - and node.node.node == nodes.Name("content", "load") - and node.node.attr in {"startswith", "endswith"} - ): - return apply_edits() - if isinstance(node, nodes.Filter) and node.name not in ( - "default items length safe string tojson trim".split() - ): - return apply_edits() - if isinstance(node, nodes.Test) and node.name not in ( - "defined false iterable mapping none string true undefined".split() - ): - return apply_edits() assignments = list(tree.find_all(nodes.Assign)) locations = [] parsed_assignments = [] diff --git a/tests/unit/test_literal_reasoning_content.py b/tests/unit/test_literal_reasoning_content.py index f9d1ec2e0..2a89ef397 100644 --- a/tests/unit/test_literal_reasoning_content.py +++ b/tests/unit/test_literal_reasoning_content.py @@ -2,7 +2,7 @@ import hashlib from pathlib import Path -from jinja2 import DictLoader, pass_context +from jinja2 import DictLoader, Environment, pass_context from jinja2.sandbox import ImmutableSandboxedEnvironment import pytest @@ -364,6 +364,59 @@ def mutate(value, count): ) +@pytest.mark.parametrize("environment", [Environment, ImmutableSandboxedEnvironment]) +@pytest.mark.parametrize( + "middle", + [ + "{{ inject() }}", + "{{ ''|inject }}", + "{% if '' is inject %}{% endif %}", + "{% include 'setup' %}", + ], +) +def test_counter_free_renderer_rejects_implicit_macro_mutation(environment, middle): + match = _QWEN_INLINE_REASONING.search(_TEMPLATE) + assert match is not None + trim = "{% set content = render_content(message.content, true)|trim %}" + source = ( + "{% macro render_content(content,count) %}{{ content }}{% endmacro %}" + + middle + + trim + + "{% if message.role == 'assistant' %}" + + match.group() + + "{% endif %}[{{ content }}]" + ) + + def render(template): + message = {"role": "assistant", "content": " answer "} + calls = [] + + @pass_context + def inject(context, value=None): + calls.append("inject") + + def replacement(content, count): + calls.append("render:" + message["role"]) + message["role"] = "user" + return content + + context["render_content"]._func = replacement + return "" + + env = environment(loader=DictLoader({"setup": "{{ inject() }}"})) + env.filters["inject"] = env.tests["inject"] = inject + output = env.from_string(template).render(message=message, inject=inject) + return output, calls, message["role"] + + expected = "[answer]", ["inject", "render:assistant"], "user" + fixed = _without_inline_reasoning_parser(source) + assert render(source) == render(source.replace(match.group(), "")) == expected + assert render(fixed) == expected + assert trim in fixed + assert not _QWEN_INLINE_REASONING.search(fixed) + assert _without_inline_reasoning_parser(fixed) == fixed + + @pytest.mark.parametrize( "middle", [ @@ -899,7 +952,7 @@ def test_plain_block_whitespace_keeps_prior_inline_parser_coverage(separator, qu @pytest.mark.parametrize("newline", ["\n", "\r\n"]) @pytest.mark.parametrize("layout", ["macros", "same_line", "branches"]) @pytest.mark.parametrize("inline_structured_reasoning", [False, True]) -def test_content_trim_is_scoped_to_the_recognized_parser( +def test_custom_macro_calls_keep_trim_while_removing_recognized_parser( preserve, newline, layout, inline_structured_reasoning ): match = _QWEN_INLINE_REASONING.search(_TEMPLATE) @@ -929,18 +982,22 @@ def test_content_trim_is_scoped_to_the_recognized_parser( fixed = chat_template_with_preserved_thinking(template) assert isinstance(fixed, str) assert preview in fixed + # Other macro calls lack a closed callable-custody proof. Retain their + # original trims while still removing only the recognized parser. + assert trim in fixed + assert not _QWEN_INLINE_REASONING.search(fixed) env = ImmutableSandboxedEnvironment(trim_blocks=True, lstrip_blocks=True) kwargs = dict( message={"role": "assistant", "content": " answer "}, preserve_thinking=preserve, ) - assert env.from_string(fixed).render(**kwargs).endswith("[ answer ]|[answer]") + assert env.from_string(fixed).render(**kwargs).endswith("[answer]|[answer]") kwargs["message"]["content"] = " beforeliteralafter " assert ( env.from_string(fixed) .render(**kwargs) .endswith( - "[ beforeliteralafter ]|[beforeliteralafter]" + "[beforeliteralafter]|[beforeliteralafter]" ) ) assert chat_template_with_preserved_thinking(fixed) == fixed @@ -1121,7 +1178,7 @@ def test_shared_content_preview_prevents_ambiguous_trim_rewrite(consumer): ) -def test_only_source_edited_parser_authorizes_its_content_trim(): +def test_custom_macro_calls_keep_trim_and_unedited_parser(): match = _QWEN_INLINE_REASONING.search(_TEMPLATE) assert match is not None parser = match.group() @@ -1142,8 +1199,10 @@ def test_only_source_edited_parser_authorizes_its_content_trim(): ) fixed = _without_inline_reasoning_parser(template) assert preview in fixed + assert answer not in fixed + assert trim in fixed env = ImmutableSandboxedEnvironment(trim_blocks=True, lstrip_blocks=True) rendered = env.from_string(fixed).render( message={"role": "assistant", "content": " HEADliteralTAIL "} ) - assert rendered == "[ HEADliteralTAIL ]|[TAIL]" + assert rendered == "[HEADliteralTAIL]|[TAIL]" From 90e36bfb36254d71294fa371156f498cb3a5f465 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 13:44:06 +0000 Subject: [PATCH 66/67] Preserve trim around renderer parameter mutations --- src/art_inference/chat_template.py | 4 +++ tests/unit/test_literal_reasoning_content.py | 36 ++++++++++++++++++++ 2 files changed, 40 insertions(+) diff --git a/src/art_inference/chat_template.py b/src/art_inference/chat_template.py index 5e1daed6d..bea7ae938 100644 --- a/src/art_inference/chat_template.py +++ b/src/art_inference/chat_template.py @@ -218,6 +218,10 @@ def apply_edits() -> str: } if any( not isinstance(node, allowed) + # Parameters and their local aliases can refer to the caller's message. + # Only the already-proven private counters may be mutated here. + or isinstance(node, nodes.NSRef) + and node.name not in counters or isinstance(node, nodes.Call) and not ( node.node == nodes.Name("raise_exception", "load") diff --git a/tests/unit/test_literal_reasoning_content.py b/tests/unit/test_literal_reasoning_content.py index 2a89ef397..c4dbb0ff1 100644 --- a/tests/unit/test_literal_reasoning_content.py +++ b/tests/unit/test_literal_reasoning_content.py @@ -364,6 +364,42 @@ def mutate(value, count): ) +@pytest.mark.parametrize( + "mutation", + [ + "{% set content.role = 'user' %}", + "{% set alias = content %}{% set alias.role = 'user' %}", + ], +) +@pytest.mark.parametrize( + "text", [" answer ", " beforeliteralafter "] +) +def test_renderer_cannot_mutate_parameter_namespace_aliases(mutation, text): + match = _QWEN_INLINE_REASONING.search(_TEMPLATE) + assert match is not None + trim = "{% set content = render_content(message.content, true)|trim %}" + source = ( + "{% macro render_content(content,count) %}" + + mutation + + "{{ content.text }}{% endmacro %}" + + "{% set message = namespace(role=message.role, text=message.content) %}" + + "{% set message.content = message %}" + + trim + + "{% if message.role == 'assistant' %}" + + match.group() + + "{% else %}[{{ content }}]{% endif %}" + ) + env = ImmutableSandboxedEnvironment() + kwargs = {"message": {"role": "assistant", "content": text}} + original = env.from_string(source).render(**kwargs) + fixed = _without_inline_reasoning_parser(source) + assert original == "[" + text.strip() + "]" + assert env.from_string(fixed).render(**kwargs) == original + assert trim in fixed + assert not _QWEN_INLINE_REASONING.search(fixed) + assert _without_inline_reasoning_parser(fixed) == fixed + + @pytest.mark.parametrize("environment", [Environment, ImmutableSandboxedEnvironment]) @pytest.mark.parametrize( "middle", From 7407b1b4588ab194669523598bdf7ebe9d9f928c Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 13:51:51 +0000 Subject: [PATCH 67/67] Keep private counters bound to validated declarations --- src/art_inference/chat_template.py | 8 +++--- tests/unit/test_literal_reasoning_content.py | 27 ++++++++++++++++++++ 2 files changed, 30 insertions(+), 5 deletions(-) diff --git a/src/art_inference/chat_template.py b/src/art_inference/chat_template.py index bea7ae938..430ee69f4 100644 --- a/src/art_inference/chat_template.py +++ b/src/art_inference/chat_template.py @@ -182,6 +182,7 @@ def apply_edits() -> str: except TemplateSyntaxError: return apply_edits() counters = global_names - {"raise_exception", "add_vision_id"} + counter_targets: set[int] = set() for name in counters: declarations = [ node @@ -194,6 +195,7 @@ def apply_edits() -> str: expected = declaration.node if len(declarations) != 1 or declarations[0].node != expected: return apply_edits() + counter_targets.add(id(declarations[0].target)) inside = {id(node) for node in renderer.find_all(nodes.Node)} if any( node.name == name @@ -241,11 +243,7 @@ def apply_edits() -> str: and node.ctx != "load" or node.name in counters and node.ctx != "load" - and not any( - node is declaration.target - for declaration in tree.body - if isinstance(declaration, nodes.Assign) - ) + and id(node) not in counter_targets for node in tree.find_all(nodes.Name) ) or any( isinstance(node, nodes.Import) diff --git a/tests/unit/test_literal_reasoning_content.py b/tests/unit/test_literal_reasoning_content.py index c4dbb0ff1..d346d25d3 100644 --- a/tests/unit/test_literal_reasoning_content.py +++ b/tests/unit/test_literal_reasoning_content.py @@ -400,6 +400,33 @@ def test_renderer_cannot_mutate_parameter_namespace_aliases(mutation, text): assert _without_inline_reasoning_parser(fixed) == fixed +@pytest.mark.parametrize("text", [" answer ", " literalafter "]) +def test_private_renderer_counter_cannot_be_rebound_to_message(text): + match = _QWEN_INLINE_REASONING.search(_TEMPLATE) + assert match is not None + trim = "{% set content = render_content(message.content, true)|trim %}" + source = ( + "{% set counter = namespace(value=0) %}" + "{% macro render_content(content,count) %}" + "{% set counter.role = 'user' %}{{ content }}{% endmacro %}" + "{% set message = namespace(role=message.role, content=message.content) %}" + "{% set counter = message %}" + + trim + + "{% if message.role == 'assistant' %}" + + match.group() + + "{% else %}[{{ content }}]{% endif %}" + ) + env = ImmutableSandboxedEnvironment() + kwargs = {"message": {"role": "assistant", "content": text}} + original = env.from_string(source).render(**kwargs) + fixed = _without_inline_reasoning_parser(source) + assert original == "[" + text.strip() + "]" + assert env.from_string(fixed).render(**kwargs) == original + assert trim in fixed + assert not _QWEN_INLINE_REASONING.search(fixed) + assert _without_inline_reasoning_parser(fixed) == fixed + + @pytest.mark.parametrize("environment", [Environment, ImmutableSandboxedEnvironment]) @pytest.mark.parametrize( "middle",