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 )