diff --git a/docs/features/additional-histories.mdx b/docs/features/additional-histories.mdx
index 24802ace7..2a31c11bf 100644
--- a/docs/features/additional-histories.mdx
+++ b/docs/features/additional-histories.mdx
@@ -33,6 +33,24 @@ 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.
+
+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.
+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:
```python
@@ -44,7 +62,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 +71,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"}
]
)
]
@@ -134,6 +152,65 @@ 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.
+
+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
+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,
+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.
+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`,
+`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..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
@@ -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)
@@ -591,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],
@@ -642,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],
@@ -809,16 +941,25 @@ 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], ...] = ()
+ validate_sources: Callable[[_SampledSourceKey | None], None] | None = None
def set(
self,
tokenized: TokenizedHistory,
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
+ self.rendered_outputs = rendered_outputs
+ self.validate_sources = _sampled_source_validator(sources)
def _fingerprint(value: object) -> str:
@@ -859,7 +1000,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")
@@ -956,6 +1107,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,
@@ -970,27 +1122,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)
@@ -1003,6 +1168,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")
@@ -2059,6 +2225,38 @@ 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):
+ 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
+
+
def _template_ids(
tokenizer: Tokenizer,
exchange: Exchange,
@@ -2106,9 +2304,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(
@@ -2749,7 +2947,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
@@ -3225,7 +3423,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
@@ -3237,7 +3439,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)
@@ -3399,7 +3601,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)
@@ -3419,18 +3625,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),
@@ -3449,8 +3658,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
]
@@ -3588,6 +3798,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 +3830,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 +3870,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 +3923,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 +3942,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,
@@ -3725,7 +3979,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
@@ -3851,6 +4105,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):
@@ -4031,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],
@@ -4039,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)
@@ -4053,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)
@@ -4118,6 +4438,182 @@ 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 = _source_exchange(source)
+ 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 _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)
+ 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 = _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)
+ 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 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
+
+
+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]],
+ 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:
+ 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,13 +4622,20 @@ 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]] = (),
+ _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]] = (
+ {} 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):
- signature = _source_signature(source)
+ signature = _source_signature(source, _fingerprints=fingerprints)
if (
message.get("role") == "assistant"
and source is not None
@@ -4145,12 +4648,22 @@ 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)
+ 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.
@@ -4223,8 +4736,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 +4747,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,
@@ -4260,10 +4771,65 @@ def _tokenize_exact_projected_chat_history(
TokenFlag.EXACT | TokenFlag.SAMPLED | TokenFlag.ASSISTANT | TokenFlag.OUTPUT
] * len(retained_ids)
logprobs[start:end] = retained_logprobs
- source_key = _sampled_source_key(source)
- if _source_stop_evidence(source, source_key)[0] == "length":
+ 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
+ ):
+ 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.
+ 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)
- next_prompt = _chat_source_prompt_tokens(sampled_sources[index + 1])
+ 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 = 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 +4839,21 @@ 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.
+ fingerprints.clear()
+ 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(
+ 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 +4872,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:
@@ -4322,10 +4904,146 @@ def _tokenize_exact_projected_chat_history(
flags=flags,
)
if _trace is not None:
- _trace.set(tokenized, source_keys, sources)
+ _trace.set(tokenized, source_keys, sources, tokenizer=tokenizer)
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
+ # 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"}:
+ return None
+ if stop == "stop" and _sampled_stop_suffix(
+ output, source=source, source_key=key, tokenizer=tokenizer
+ ):
+ continue
+ try:
+ 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,
+ # 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]
+ 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] = []
+ 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 +5250,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 +5258,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 = (
@@ -4625,12 +5293,14 @@ 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)
+ original_template = template
+ template, defaults = _resolved_chat_template(
+ resolved_tokenizer, template, history.tools
+ )
kwargs = {
- **default_chat_template_kwargs_for_template(template),
+ **defaults,
**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 +5346,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(
@@ -4937,6 +5618,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
@@ -4959,6 +5641,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(
@@ -4973,22 +5721,48 @@ def source_matches_context(source: object) -> bool:
canonical_assistant_mask,
direct_bounds or None,
)
- assistant_mask = _translate_token_mask(
- canonical_rendered,
- rendered,
- canonical_assistant_mask,
- tokenizer=resolved_tokenizer,
- )
- output_mask = _translate_token_mask(
- canonical_rendered,
- rendered,
- canonical_output_mask,
- tokenizer=resolved_tokenizer,
- )
- stop_mask = _translate_token_mask(canonical_rendered, rendered, canonical_stop_mask)
- length_stop_mask = _translate_token_mask(
- canonical_rendered, rendered, canonical_length_stop_mask
- )
+ 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)
@@ -5472,6 +6246,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"
@@ -5578,7 +6355,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(
@@ -5589,6 +6366,14 @@ 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,
+ )
):
return exact
@@ -6137,6 +6922,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,
@@ -6220,6 +7006,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
@@ -6295,7 +7085,13 @@ 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),
+ tokenizer=resolved_tokenizer,
+ )
return tokenized
@@ -6373,7 +7169,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
@@ -6563,7 +7359,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
@@ -6655,7 +7451,9 @@ 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,
+ _copied_context: bool = False,
) -> TokenizedHistory:
if isinstance(history, LegacyHistory):
if model is None:
@@ -6677,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)
@@ -6690,11 +7496,17 @@ 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.
+ 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
)
- has_length_stop = can_render and _history_has_length_stop(history)
- needs_synthetic_stop = _history_needs_synthetic_stop(history, tokenizer)
needs_render = (
render_state.needs_render
or override_requires_render
@@ -6703,12 +7515,43 @@ 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):
+ 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
+ (
+ 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_request_roles
and not needs_synthetic_stop
and not override_requires_render
and not render_state.context_changed
@@ -6721,6 +7564,9 @@ def _tokenize_history(
_projection_validated or render_state.projection_matches is True
),
_trace=_trace,
+ _strict_sources=True,
+ _prior=_prior,
+ _fingerprints=fingerprints,
)
)
):
@@ -6734,9 +7580,18 @@ 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 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)
+ ),
_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
@@ -6754,18 +7609,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(),
@@ -6805,8 +7664,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 +7716,23 @@ 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,
+ _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,
+ 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(
@@ -6848,6 +7764,50 @@ 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:
+ """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
+ ):
+ assert builder.validate_sources is not None
+ builder.validate_sources(None)
+ _mark_sampled_stops(
+ value.tokens,
+ value.flags,
+ builder.trace.source_keys,
+ builder.trace.sources,
+ tokenizer=tokenizer,
+ )
+ _validate_completed_sources(builders)
+
+
def tokenize_trajectory(
trajectory: Trajectory,
*,
@@ -6877,8 +7837,14 @@ 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 = []
+ stop_builders: list[_TraceBuilder | None] = []
+ for history, copied in zip(histories, context_sources, strict=True):
+ trace = _TraceBuilder() if len(histories) > 1 else None
+ result = tokenize_history(
history,
model=model if isinstance(history, LegacyHistory) else history.model,
base_model=base_model,
@@ -6886,9 +7852,15 @@ 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)
+ 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)
if not multi_history:
return _materialize_trajectory(tokenized[0], trajectory)
return TokenizedMultiHistoryTrajectory(
@@ -6914,6 +7886,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(
@@ -6929,11 +7902,14 @@ 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")
tokenized_histories.append(tokenized)
traces.append(trace_builder.trace)
+ builders.append(trace_builder)
+ _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 20dd21c7e..93b18f294 100644
--- a/src/art_inference/chat_template.py
+++ b/src/art_inference/chat_template.py
@@ -25,10 +25,125 @@
"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_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 _QWEN_INLINE_STATEMENTS
+ )
+)
+
+
+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.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()
+
+ def operations(text: str):
+ return WithoutWhitespace().visit(env.parse(text)).body
+
+ operation = operations(
+ "".join("{% " + statement + " %}" for statement in _QWEN_INLINE_STATEMENTS)
+ )
+ 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")
+ ] + [len(template)]
+ blocks: list[tuple[int, int, int, int]] = []
+ cursor = 0
+ opening = None
+ try:
+ 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":
+ 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: 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 operations(template[start:end]) == operation:
+ # 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:
+ return template
+ # Dropping structured reasoning must not trim the visible assistant body.
+ 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
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 +151,7 @@ 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)
replacements = (
(
_QWEN_DROP_PRIOR_THINKING,
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:
diff --git a/tests/unit/test_literal_reasoning_content.py b/tests/unit/test_literal_reasoning_content.py
new file mode 100644
index 000000000..1c8056ec8
--- /dev/null
+++ b/tests/unit/test_literal_reasoning_content.py
@@ -0,0 +1,421 @@
+from copy import deepcopy
+import hashlib
+from pathlib import Path
+
+from jinja2.sandbox import ImmutableSandboxedEnvironment
+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,
+)
+
+_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",
+ "",
+)
+
+
+@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)
+
+ 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
+
+
+@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
+
+ 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
+
+
+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,
+ {"role": "assistant", "content": "answer", "reasoning_content": "prior reason"},
+ _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
+
+ 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
+ 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
+
+ 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)
+ assert fixed == _FIXED + literal
+ assert "headliteraltail" in _render(
+ 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
+
+
+@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
+
+
+@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_evidence_reuse.py b/tests/unit/trajectories/test_evidence_reuse.py
new file mode 100644
index 000000000..31ff8d55c
--- /dev/null
+++ b/tests/unit/trajectories/test_evidence_reuse.py
@@ -0,0 +1,511 @@
+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) == 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) == 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
+
+ monkeypatch.setattr(module, "_WARNED_PREFIX_RETOKENIZATION", False)
+ 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(
+ "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
+
+
+@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"
+ 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_literal_thinking_off.py b/tests/unit/trajectories/test_literal_thinking_off.py
index 16188c733..f956c0b09 100644
--- a/tests/unit/trajectories/test_literal_thinking_off.py
+++ b/tests/unit/trajectories/test_literal_thinking_off.py
@@ -12,12 +12,14 @@
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 = (
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:
@@ -125,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]
@@ -137,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, "_preserve_literal_thinking_off_content", lambda *args: None
- )
- _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
]
@@ -162,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]["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
@@ -183,9 +198,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"
@@ -222,25 +235,31 @@ 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, 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
+ 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
@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)
@@ -249,7 +268,10 @@ def test_explicit_empty_reasoning_is_preserved(field: str) -> None:
)
== _LITERAL
)
- assert tokenizer.calls[0][-1]["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
@@ -278,7 +300,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
]
@@ -403,21 +426,21 @@ 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
+ _tokenize, "chat_template_with_preserved_thinking", lambda value: value
)
_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"])]
@@ -441,3 +464,189 @@ 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]
+
+
+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
diff --git a/tests/unit/trajectories/test_native_terminal.py b/tests/unit/trajectories/test_native_terminal.py
new file mode 100644
index 000000000..da2339027
--- /dev/null
+++ b/tests/unit/trajectories/test_native_terminal.py
@@ -0,0 +1,253 @@
+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
+ )
+
+
+@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=finish, sampled_eos=sampled_eos, 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)
+ 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)]
+ 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]
+ 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(
+ 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
new file mode 100644
index 000000000..0e4af63c9
--- /dev/null
+++ b/tests/unit/trajectories/test_recorded_boundaries.py
@@ -0,0 +1,925 @@
+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 - (
+ tool_position == 1
+ )
+ 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:
+ # 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"])
+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
+
+
+@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_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
+ )
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..fdfd2cf90
--- /dev/null
+++ b/tests/unit/trajectories/test_resolved_stop_authority.py
@@ -0,0 +1,400 @@
+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)
+ 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
+
+
+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)
+
+
+@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(
+ 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", 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 [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] == [
+ ord("r") + 100,
+ 9,
+ ord("r") + 100,
+ 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, eos=8), 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_historical_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.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)
+ 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 0e04a5268..fc86d06c3 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,
]
@@ -766,6 +765,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)
@@ -821,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(
@@ -834,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
@@ -931,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)
@@ -1380,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 = ""
@@ -1388,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,
]
@@ -3698,19 +3798,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,
@@ -4624,12 +4737,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 +7330,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 +7502,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(
@@ -8964,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
)