From cdeb6e2a916bbc7e6b3f725ad5b697559c03404b Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 05:21:45 +0000 Subject: [PATCH 1/2] Reuse exact chat-template prefixes during long-history tokenization --- src/art/trajectories/_render_cache.py | 210 ++++++++++ src/art/trajectories/_tokenize.py | 118 +++++- .../trajectories/test_prefix_render_cache.py | 368 ++++++++++++++++++ .../test_render_cache_eligibility.py | 150 +++++++ 4 files changed, 838 insertions(+), 8 deletions(-) create mode 100644 src/art/trajectories/_render_cache.py create mode 100644 tests/unit/trajectories/test_prefix_render_cache.py create mode 100644 tests/unit/trajectories/test_render_cache_eligibility.py diff --git a/src/art/trajectories/_render_cache.py b/src/art/trajectories/_render_cache.py new file mode 100644 index 000000000..db976f1d0 --- /dev/null +++ b/src/art/trajectories/_render_cache.py @@ -0,0 +1,210 @@ +"""Conservative eligibility for invocation-local chat-render reuse.""" + +import inspect +import math +import sys +from types import CodeType +from typing import cast + + +def _render_context_key(value: object) -> object: + """Snapshot plain JSON without losing mapping order or scalar types.""" + kind = type(value) + if kind in (str, int, bool, type(None)): + return kind, value + if kind is float and math.isfinite(cast(float, value)): + return kind, repr(value) + if kind is list: + return kind, tuple(_render_context_key(item) for item in cast(list, value)) + if kind is dict and all(type(key) is str for key in cast(dict, value)): + return kind, tuple( + (key, _render_context_key(item)) for key, item in cast(dict, value).items() + ) + raise TypeError("Not a plain JSON rendering context") + + +def _code_contains(code, candidate, name): + return (code is candidate and code.co_name == name) or any( + _code_contains(child, candidate, name) + for child in code.co_consts + if isinstance(child, CodeType) + ) + + +def cacheable_chat_template(tokenizer, template, tools, kwargs, messages) -> bool: + """Admit a finite, nonmutating Jinja subset, never arbitrary renderer purity. + + Callers must key all effective context and keep this cache local to one + tokenization. Unsupported syntax, context or overrides retain normal rendering. + No Transformers import or global eligibility memo is performed here. + """ + try: + base_module = sys.modules.get("transformers.tokenization_utils_base") + chat = sys.modules.get("transformers.utils.chat_template_utils") + if base_module is None or chat is None or type(template) is not str: + return False + base = base_module.PreTrainedTokenizerBase + cls = type(tokenizer) + module = sys.modules.get(cls.__module__) + if ( + not isinstance(tokenizer, base) + or not cls.__module__.startswith("transformers.") + or getattr(module, cls.__name__, None) is not cls + or type(tokenizer.chat_template) not in (str, type(None)) + or inspect.getattr_static(cls, "special_tokens_map") + is not inspect.getattr_static(base, "special_tokens_map") + ): + return False + for name in ("apply_chat_template", "get_chat_template"): + method = getattr(tokenizer, name) + if method.__self__ is not tokenizer or method.__func__ is not getattr( + base, name + ): + return False + # Check before the stock property calls str(): custom objects can hide + # mutable state even if their resulting special-token strings look plain. + added_token = sys.modules["tokenizers"].AddedToken + special = tokenizer._special_tokens_map + if type(special) is not dict or any( + type(key) is not str or type(value) not in (str, type(None), added_token) + for key, value in special.items() + ): + return False + _render_context_key([messages, tools, kwargs, tokenizer.special_tokens_map]) + if type(kwargs) is not dict or kwargs.get("continue_final_message"): + return False + + from jinja2 import defaults, nodes + from jinja2.runtime import LoopContext + from jinja2.sandbox import ImmutableSandboxedEnvironment + from jinja2.utils import Namespace + + compiled = chat._compile_jinja_template(template) + env = compiled.environment + if type(env) is not ImmutableSandboxedEnvironment: + return False + tree = env.parse(template) # HF's environment understands {% generation %}. + parents = { + child: node + for node in (tree, *tree.find_all(nodes.Node)) + for child in node.iter_child_nodes() + } + macros = {node.name for node in tree.find_all(nodes.Macro)} + filters = set("default length tojson trim items string safe".split()) + tests = set("string iterable mapping none undefined true false defined".split()) + methods = set( + "get items keys values startswith endswith strip lstrip rstrip split rsplit replace lower upper join".split() + ) + # A method object can expose an address when printed or aliased. Permit + # only direct calls of the explicitly nonmutating methods above. + method_names = { + name + for cls in (str, dict, list, tuple, int, float, LoopContext) + for name in dir(cls) + if callable(getattr(cls, name)) + } + structural = set( + "Template Output TemplateData Const Name Getattr Getitem Slice If For Assign AssignBlock NSRef Macro Call CallBlock Keyword Filter Test Compare Operand And Or Not Neg Pos Add Sub Mul Div FloorDiv Mod Pow Concat List Tuple Dict Pair CondExpr Break Continue ExtensionAttribute".split() + ) + compiler = getattr( + chat, + "_cached_compile_jinja_template", + chat._compile_jinja_template.__wrapped__, + ) + for node in (tree, *tree.find_all(nodes.Node)): + parent = parents.get(node) + direct_call = isinstance(parent, nodes.Call) and parent.node is node + if type(node).__name__ not in structural: + return False + if isinstance(node, nodes.Name): + if node.name in {"self", "super", "caller"}: + return False + if node.name in env.globals or node.name in macros: + if not direct_call: + return False + if isinstance(node, nodes.Getattr): + if node.attr.startswith("_") or ( + node.attr in method_names and not direct_call + ): + return False + if isinstance(node, nodes.ExtensionAttribute) and not direct_call: + return False + if isinstance(node, nodes.Getitem): + # Dynamic lookup can fetch a callable or renderer/context object. + parts = ( + (node.arg.start, node.arg.stop, node.arg.step) + if isinstance(node.arg, nodes.Slice) + else (node.arg,) + ) + parts = tuple( + part.node if isinstance(part, (nodes.Neg, nodes.Pos)) else part + for part in parts + ) + if any( + part is not None + and not (isinstance(part, nodes.Const) and type(part.value) is int) + for part in parts + ): + return False + if isinstance(node, nodes.Call): + target = node.node + if node.dyn_args is not None or node.dyn_kwargs is not None: + return False + if isinstance(target, nodes.Name): + if target.name in macros: + continue + helper = env.globals.get(target.name) + if target.name == "namespace" and helper is Namespace: + continue + if target.name == "raise_exception" and _code_contains( + compiler.__code__, + getattr(helper, "__code__", None), + "raise_exception", + ): + continue + return False + if isinstance(target, nodes.Getattr) and target.attr in methods: + continue + if isinstance(target, nodes.ExtensionAttribute) and isinstance( + parent, nodes.CallBlock + ): + extension = env.extensions.get(target.identifier) + helper = getattr( + getattr(extension, target.name, None), "__func__", None + ) + if target.name == "_generation_support" and _code_contains( + compiler.__code__, + getattr(helper, "__code__", None), + "_generation_support", + ): + continue + return False + if isinstance(node, (nodes.Filter, nodes.Test)): + if isinstance(node, nodes.Test): + if ( + node.name not in tests + or env.tests.get(node.name) + is not defaults.DEFAULT_TESTS[node.name] + ): + return False + elif node.name not in filters: + return False + elif node.name == "tojson": + if not _code_contains( + compiler.__code__, + getattr(env.filters.get(node.name), "__code__", None), + "tojson", + ): + return False + elif ( + env.filters.get(node.name) + is not defaults.DEFAULT_FILTERS[node.name] + ): + return False + if node.name == "items" and not ( + isinstance(parent, nodes.For) and parent.iter is node + ): + return False # Otherwise a generator's repr can expose identity. + return True + except Exception: + return False diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index 9b692894c..3a9f2b6c6 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -59,6 +59,7 @@ ) from ._history import _model_matches from ._protocols import Exchange +from ._render_cache import _render_context_key, cacheable_chat_template _TOKEN_ID = re.compile(r"token_id:(\d+)$") _WARNED_PREFIX_RETOKENIZATION = False @@ -92,6 +93,73 @@ def __call__( ) -> str: ... +class _PrefixChatRenderCache: + """Reuse exact baseline prefixes within one history's rendering context. + + Probes still locate every assistant span. Equal completed text does not prove + equal generation prompts, so changed prefixes never reuse baseline renders. + Keep one baseline plus bounded prefix deltas, not quadratic rendered strings. + """ + + _MAX_BYTES = 8 * 1024 * 1024 + _MAX_ENTRIES = 1024 + + def __init__(self, render: _ChatRender) -> None: + self.render = render + self.context: tuple[object, ...] | None = None + self.settings: object = None + self.text = "" + self.prefixes: dict[tuple[int, bool], tuple[int, str]] = {} + self.bytes = 0 + + def for_messages( + self, messages: list[dict[str, Any]], text: str, *, settings: object = None + ) -> _ChatRender: + try: + context = tuple(_render_context_key(message) for message in messages) + except (TypeError, RecursionError): + return self.render + if self.context is None or settings != self.settings: + self.context, self.text = context, text + self.settings = settings + self.prefixes.clear() + self.bytes = 0 + common = 0 + for original, current in zip(self.context, context): + if original != current: + break + common += 1 + + def render( + selected_messages: list[dict[str, Any]], *, add_generation_prompt: bool + ) -> str: + count = len(selected_messages) + if count > common or any( + original is not current + for original, current in zip(messages, selected_messages) + ): + return self.render( + selected_messages, add_generation_prompt=add_generation_prompt + ) + key = count, add_generation_prompt + if key in self.prefixes: + prefix, tail = self.prefixes[key] + return self.text[:prefix] + tail + value = self.render( + selected_messages, add_generation_prompt=add_generation_prompt + ) + if len(self.prefixes) < self._MAX_ENTRIES: + prefix = _common_prefix_length(self.text, value) + tail = value[prefix:] + size = 256 + 4 * len(tail) + if self.bytes + size <= self._MAX_BYTES: + self.prefixes[key] = prefix, tail + self.bytes += size + return value + + return render + + class _TokenChatRender(Protocol): def __call__( self, @@ -4530,13 +4598,11 @@ def raw_render( ) ) - def render_text( + def render_normalized_text( selected_messages: list[dict[str, Any]], *, add_generation_prompt: bool ) -> str: value = resolved_tokenizer.apply_chat_template( - normalize_tool_call_arguments_for_chat_template( - selected_messages, template - ), + selected_messages, tools=history.tools, tokenize=False, add_generation_prompt=add_generation_prompt, @@ -4547,17 +4613,53 @@ def render_text( raise TypeError("Chat template did not render text") return value + def render_text( + selected_messages: list[dict[str, Any]], *, add_generation_prompt: bool + ) -> str: + return render_normalized_text( + normalize_tool_call_arguments_for_chat_template( + selected_messages, template + ), + add_generation_prompt=add_generation_prompt, + ) + + prefix_render_cache = _PrefixChatRenderCache(render_normalized_text) + def segmented_render( selected_messages: list[dict[str, Any]], *, add_generation_prompt: bool ) -> tuple[list[int], list[bool]]: try: - text = render_text( - selected_messages, add_generation_prompt=add_generation_prompt - ) + if cacheable_chat_template( + resolved_tokenizer, template, history.tools, kwargs, selected_messages + ): + # Normalization is message-local. Only the admitted nonmutating + # renderer may share its normalized messages between prefixes. + selected_messages = normalize_tool_call_arguments_for_chat_template( + selected_messages, template + ) + text = render_normalized_text( + selected_messages, add_generation_prompt=add_generation_prompt + ) + span_render = prefix_render_cache.for_messages( + selected_messages, + text, + settings=_render_context_key( + [ + history.tools, + kwargs, + getattr(resolved_tokenizer, "special_tokens_map"), + ] + ), + ) + else: + text = render_text( + selected_messages, add_generation_prompt=add_generation_prompt + ) + span_render = render_text spans = _assistant_char_spans( selected_messages, text, - render_text, + span_render, add_generation_prompt=add_generation_prompt, ) except (TypeError, KeyError): diff --git a/tests/unit/trajectories/test_prefix_render_cache.py b/tests/unit/trajectories/test_prefix_render_cache.py new file mode 100644 index 000000000..63f282103 --- /dev/null +++ b/tests/unit/trajectories/test_prefix_render_cache.py @@ -0,0 +1,368 @@ +from copy import deepcopy +from datetime import datetime, timedelta +import json +import string + +from openai.types.chat import ChatCompletion +import pytest + +import art.trajectories as tr +from art.trajectories import _tokenize as tokenization + + +@pytest.fixture(autouse=True) +def restore_retokenization_warning(monkeypatch): + monkeypatch.setattr(tokenization, "_WARNED_PREFIX_RETOKENIZATION", False) + + +def test_prefix_cache_preserves_order_types_generation_and_probe_context(): + calls = [] + + def render(messages, *, add_generation_prompt): + calls.append(deepcopy(messages)) + return json.dumps(messages) + str(add_generation_prompt) + + messages = [{"role": "user", "content": "snow雪"}, {"a": -0.0, "b": True}] + cache = tokenization._PrefixChatRenderCache(render) + expected = render(messages, add_generation_prompt=False) + cached = cache.for_messages(messages, expected) + for generation in (True, False): + assert cached(messages[:1], add_generation_prompt=generation) == render( + messages[:1], add_generation_prompt=generation + ) + before = len(calls) + assert cached(messages[:1], add_generation_prompt=True).endswith("True") + assert len(calls) == before + + # Mutating a new probe, reordering mapping keys, and scalar equality must + # never alias the baseline context, including -0.0 versus +0.0. + for replacement in ( + {"a": 0.0, "b": True}, + {"b": True, "a": -0.0}, + {"a": -0.0, "b": 1}, + ): + probe = [deepcopy(messages[0]), replacement] + current = cache.for_messages(probe, render(probe, add_generation_prompt=False)) + assert current(probe, add_generation_prompt=False) == render( + probe, add_generation_prompt=False + ) + reordered = list(reversed(messages)) + assert cached(reordered, add_generation_prompt=False) == render( + reordered, add_generation_prompt=False + ) + + +def test_prefix_cache_bounds_storage_and_does_not_cache_failures(monkeypatch): + monkeypatch.setattr(tokenization._PrefixChatRenderCache, "_MAX_BYTES", 520) + monkeypatch.setattr(tokenization._PrefixChatRenderCache, "_MAX_ENTRIES", 2) + calls = 0 + + def render(messages, *, add_generation_prompt): + nonlocal calls + calls += 1 + if not messages: + raise ValueError("empty history") + return "🙂" * len(messages) + ("?" if add_generation_prompt else "") + + messages = [{"content": str(i)} for i in range(20)] + cache = tokenization._PrefixChatRenderCache(render) + cached = cache.for_messages(messages, render(messages, add_generation_prompt=False)) + for _ in range(2): + with pytest.raises(ValueError, match="empty history"): + cached([], add_generation_prompt=True) + for i in range(1, 21): + assert cached(messages[:i], add_generation_prompt=True) == "🙂" * i + "?" + assert len(cache.prefixes) <= 2 + assert cache.bytes <= 520 + assert calls > 20 # Full-cache misses and exceptions continue to render. + + +def test_cache_settings_and_mutated_messages_invalidate_previous_prefixes(): + settings = {"tool": "lookup"} + + def render(messages, *, add_generation_prompt): + return settings["tool"] + json.dumps(messages) + str(add_generation_prompt) + + messages = [{"role": "user", "content": "first"}] + cache = tokenization._PrefixChatRenderCache(render) + for tool, content in ( + ("lookup", "first"), + ("other", "first"), + ("other", "changed"), + ): + settings["tool"], messages[0]["content"] = tool, content + expected = render(messages, add_generation_prompt=True) + current = cache.for_messages( + messages, expected, settings=tokenization._render_context_key(settings) + ) + assert current(messages, add_generation_prompt=True) == expected + + +def test_prefix_deltas_do_not_retain_quadratic_text(): + def render(messages, *, add_generation_prompt): + return "".join(m["content"] for m in messages) + ( + "?" if add_generation_prompt else "" + ) + + messages = [{"content": "雪" * 4096} for _ in range(128)] + cache = tokenization._PrefixChatRenderCache(render) + cached = cache.for_messages(messages, render(messages, add_generation_prompt=False)) + for i in range(1, 129): + assert ( + cached(messages[:i], add_generation_prompt=True) == "雪" * (4096 * i) + "?" + ) + assert len(cache.prefixes) == 128 + assert sum(len(tail) for _, tail in cache.prefixes.values()) == 128 + + +@pytest.mark.parametrize("value", [object(), float("nan"), float("inf"), {1: "x"}]) +def test_prefix_cache_bypasses_non_json_context(value): + def render(messages, *, add_generation_prompt): + return "unchanged" + + cache = tokenization._PrefixChatRenderCache(render) + assert cache.for_messages([{"content": value}], "unchanged") is render + + +def test_later_generation_split_is_recomputed_when_completed_suffix_is_equal(): + messages = [ + {"role": "user", "content": "x"}, + {"role": "assistant", "content": "A"}, + {"role": "user", "content": "y"}, + {"role": "assistant", "content": "HELLO"}, + ] + + def render(selected, *, add_generation_prompt): + text = "".join(f"<{m['role']}>{m['content']}!" for m in selected) + if add_generation_prompt: + text += "" + if any(m["content"] == "B" for m in selected): + text += "H" + return text + + cache = tokenization._PrefixChatRenderCache(render) + original = render(messages, add_generation_prompt=False) + tokenization._assistant_char_spans( + messages, + original, + cache.for_messages(messages, original), + add_generation_prompt=False, + ) + probe = deepcopy(messages) + probe[1]["content"] = "B" + text = render(probe, add_generation_prompt=False) + assert text == original.replace(">A!", ">B!") + actual = tokenization._assistant_char_spans( + probe, text, cache.for_messages(probe, text), add_generation_prompt=False + ) + expected = tokenization._assistant_char_spans( + probe, text, render, add_generation_prompt=False + ) + assert actual == expected + start, end = actual[-1] + assert text[start:end] == "ELLO!" + + +_TEMPLATE = """{% for message in messages %}<{{ message.role }}>{{ message.content or '' }} +{% if message.reasoning_content %}{{ message.reasoning_content }}{% endif %} +{% for call in message.tool_calls or [] %}{% set tool_call = call.function %}{{ tool_call.name }}({% for k, v in tool_call.arguments.items() %}{{ k }}={{ v|tojson }};{% endfor %}){% endfor %} +{% if message.role == 'assistant' %}${% else %}{% endif %}{% endfor %} +{% if add_generation_prompt %}{% endif %}""" + + +def _tokenizer(template): + tokenizers = pytest.importorskip("tokenizers") + transformers = pytest.importorskip("transformers") + vocab = {char: i for i, char in enumerate(dict.fromkeys(string.printable + "雪🙂"))} + vocab["[UNK]"] = len(vocab) + backend = tokenizers.Tokenizer( + tokenizers.models.WordLevel(vocab, unk_token="[UNK]") + ) + backend.pre_tokenizer = tokenizers.pre_tokenizers.Split("", behavior="isolated") + return transformers.PreTrainedTokenizerFast( + tokenizer_object=backend, + eos_token="$", + unk_token="[UNK]", + chat_template=template, + ) + + +def _history(turns, *, reasoning=False, refusal=False, length=False, tokenizer=None): + messages = [{"role": "user", "content": "snow雪🙂"}] + exchanges = [] + for i in range(turns): + message = { + "role": "assistant", + "content": " answer ", + "tool_calls": [ + { + "id": f"call-{i}", + "type": "function", + "function": {"name": "lookup", "arguments": json.dumps({"i": i})}, + } + ], + } + if reasoning: + message["reasoning"] = "consider 雪" + if refusal: + message["refusal"] = "cannot do that" + choice = { + "index": 0, + "message": message, + "finish_reason": "length" if length and i == turns - 1 else "tool_calls", + } + if tokenizer is not None: + + def tokens(selected, generation): + return tokenizer.apply_chat_template( + tokenization.normalize_tool_call_arguments_for_chat_template( + selected, _TEMPLATE + ), + tokenize=True, + add_generation_prompt=generation, + return_dict=False, + ) + + prompt, completed = ( + tokens(messages, True), + tokens([*messages, message], False), + ) + assert completed[: len(prompt)] == prompt + output = completed[len(prompt) :] + choice.update( + prompt_token_ids=prompt, + token_ids=output, + logprobs={ + "content": [ + { + "token": f"token_id:{token}", + "logprob": -0.1, + "bytes": [], + "top_logprobs": [], + } + for token in output + ] + }, + ) + response = ChatCompletion.model_validate( + { + "id": f"response-{i}", + "object": "chat.completion", + "created": i, + "model": "test/model", + "choices": [choice], + } + ) + start = datetime(2026, 1, 1) + timedelta(seconds=i) + exchanges.append( + tr.ChatCompletionsExchange( + request=tr.ChatCompletionsRequest( + model="test/model", messages=deepcopy(messages) + ), + response=response, + start_time=start, + end_time=start + timedelta(milliseconds=1), + ) + ) + messages.extend( + [message, {"role": "tool", "tool_call_id": f"call-{i}", "content": "ok"}] + ) + return tr.Trajectory( + exchanges=tr.TrajectoryExchanges(chat_completions=exchanges) + ).chat_completions_history() + + +@pytest.mark.parametrize( + "reasoning,refusal,length", + [ + (False, False, False), + (True, False, False), + (False, True, False), + (True, True, True), + ], +) +@pytest.mark.parametrize("ends_with_assistant", [False, True]) +def test_cached_tool_probes_match_all_tokenized_fields( + monkeypatch, reasoning, refusal, length, ends_with_assistant +): + history = _history(6, reasoning=reasoning, refusal=refusal, length=length) + if not ends_with_assistant: + history.messages.append( + {"role": "tool", "tool_call_id": "call-5", "content": "ok"} + ) + history.message_sources.append(None) + tokenizer = _tokenizer(_TEMPLATE) + trace = tokenization._TraceBuilder() + actual = tokenization._tokenize_chat_view( + history, + tokenizer=tokenizer, + base_model=None, + chat_template=_TEMPLATE, + chat_template_kwargs=None, + _trace=trace, + ) + monkeypatch.setattr( + tokenization, + "cacheable_chat_template", + lambda *args: False, + ) + expected_trace = tokenization._TraceBuilder() + expected = tokenization._tokenize_chat_view( + history, + tokenizer=tokenizer, + base_model=None, + chat_template=_TEMPLATE, + chat_template_kwargs=None, + _trace=expected_trace, + ) + assert actual.model_dump_json() == expected.model_dump_json() + assert trace.trace is not None and expected_trace.trace is not None + assert trace.trace.source_keys == expected_trace.trace.source_keys + assert trace.trace.sources == expected_trace.trace.sources + assert any(f & tr.TokenFlag.OUTPUT for f in actual.flags) + + +def test_cached_probes_preserve_native_logprobs_stop_and_source_bindings(monkeypatch): + tokenizer = _tokenizer(_TEMPLATE) + history = _history(4, tokenizer=tokenizer) + actual = history.tokenize(tokenizer=tokenizer, chat_template=_TEMPLATE) + monkeypatch.setattr( + tokenization, + "cacheable_chat_template", + lambda *args: False, + ) + expected = history.tokenize(tokenizer=tokenizer, chat_template=_TEMPLATE) + assert actual.model_dump_json() == expected.model_dump_json() + assert -0.1 in actual.logprobs + assert any(flag & tr.TokenFlag.STOP for flag in actual.flags) + assert any(flag & tr.TokenFlag.SAMPLED for flag in actual.flags) + + +@pytest.mark.parametrize("turns", [4, 8, 16]) +def test_tool_probe_scaling_reuses_only_unchanged_prefixes(monkeypatch, turns): + from transformers import tokenization_utils_base + + history = _history(turns) + tokenizer = _tokenizer(_TEMPLATE) + original = tokenization_utils_base.render_jinja_template + calls = 0 + + def counted(*args, **kwargs): + nonlocal calls + calls += 1 + return original(*args, **kwargs) + + monkeypatch.setattr(tokenization_utils_base, "render_jinja_template", counted) + actual = history.tokenize(tokenizer=tokenizer, chat_template=_TEMPLATE) + cached_calls = calls + calls = 0 + monkeypatch.setattr( + tokenization, + "cacheable_chat_template", + lambda *args: False, + ) + expected = history.tokenize(tokenizer=tokenizer, chat_template=_TEMPLATE) + assert actual.model_dump_json() == expected.model_dump_json() + # This removes repeated unchanged prefixes, not the remaining quadratic + # changed-prefix work. A zero-use cache must not pass this regression. + assert cached_calls <= calls - turns * (turns - 1) diff --git a/tests/unit/trajectories/test_render_cache_eligibility.py b/tests/unit/trajectories/test_render_cache_eligibility.py new file mode 100644 index 000000000..c9f4504bd --- /dev/null +++ b/tests/unit/trajectories/test_render_cache_eligibility.py @@ -0,0 +1,150 @@ +from types import MethodType + +import pytest + +from art.trajectories._render_cache import cacheable_chat_template + + +@pytest.fixture +def tokenizer(): + tokenizers = pytest.importorskip("tokenizers") + transformers = pytest.importorskip("transformers") + backend = tokenizers.Tokenizer( + tokenizers.models.WordLevel({"[UNK]": 0}, unk_token="[UNK]") + ) + return transformers.PreTrainedTokenizerFast( + tokenizer_object=backend, unk_token="[UNK]", chat_template="{{ messages }}" + ) + + +def eligible(tokenizer, template, **context): + return cacheable_chat_template( + tokenizer, + template, + context.get("tools"), + context.get("kwargs", {}), + context.get("messages", [{"role": "user", "content": "hello"}]), + ) + + +def test_stock_render_with_local_macro_namespace_and_generation(tokenizer): + template = """{% set ns = namespace(n=0) %} +{% macro content(message) %}{{ message.content|trim }}{% endmacro %} +{% for message in messages[::-1] %}{% set ns.n = ns.n + 1 %} +{% generation %}{{ content(message) }}{% endgeneration %}{% endfor %} +{% for key, value in tools[0]|items %}{{ key }}={{ value|tojson }}{% endfor %} +{{ ns.n }}{% if add_generation_prompt %}assistant{% endif %}""" + tools = [{"name": "lookup"}] + assert eligible(tokenizer, template, tools=tools) + assert "hello" in tokenizer.apply_chat_template( + [{"role": "user", "content": "hello"}], + tools=tools, + chat_template=template, + tokenize=False, + ) + + +@pytest.mark.parametrize( + "template", + [ + "{{ strftime_now('%s') }}", + "{{ lipsum() }}", + "{{ messages|random }}", + "{{ messages|map('random')|list }}", + "{{ messages|attr('clear')() }}", + "{% set f = namespace %}{{ f() }}", + "{% set f = messages.clear %}{{ f() }}", + "{{ messages.clear() }}", + "{{ messages[0].get }}", + "{{ messages[0]|items }}", + "{% for message in messages %}{{ loop.cycle }}{% endfor %}", + "{{ messages[0]['get']('role') }}", + "{{ messages[0] is sameas messages[1] }}", + "{{ self }}", + "{% include 'other' %}", + "{% invalid %}", + ], +) +def test_unsupported_or_nondeterministic_syntax_bypasses(tokenizer, template): + assert not eligible(tokenizer, template) + + +@pytest.mark.parametrize( + "value", [object(), float("nan"), float("inf"), {1: "x"}, (1,)] +) +@pytest.mark.parametrize("field", ["messages", "tools", "kwargs"]) +def test_non_plain_context_bypasses(tokenizer, field, value): + assert not eligible(tokenizer, "{{ messages }}", **{field: {"value": value}}) + + +def test_mutable_context_is_rechecked_and_cycles_bypass(tokenizer): + tools = [{"name": "lookup"}] + assert eligible(tokenizer, "{{ tools|tojson }}", tools=tools) + tools[0]["callback"] = lambda: None + assert not eligible(tokenizer, "{{ tools|tojson }}", tools=tools) + cycle = [] + cycle.append(cycle) + assert not eligible(tokenizer, "{{ messages }}", messages=cycle) + + +@pytest.mark.parametrize("method", ["apply_chat_template", "get_chat_template"]) +def test_custom_method_is_never_called(tokenizer, monkeypatch, method): + def mutating(self, *args, **kwargs): + raise AssertionError("eligibility must not invoke custom renderers") + + monkeypatch.setattr(tokenizer, method, MethodType(mutating, tokenizer)) + assert not eligible(tokenizer, "{{ messages }}") + + +def test_custom_class_named_template_and_special_objects_bypass(tokenizer, monkeypatch): + class Custom(type(tokenizer)): + pass + + monkeypatch.setattr(tokenizer, "__class__", Custom) + assert not eligible(tokenizer, "{{ messages }}") + monkeypatch.undo() + monkeypatch.setattr(tokenizer, "chat_template", {"named": "{{ messages }}"}) + assert not eligible(tokenizer, "named") + monkeypatch.undo() + + class Token: + def __str__(self): + raise AssertionError("do not stringify custom special tokens") + + monkeypatch.setitem(tokenizer._special_tokens_map, "eos_token", Token()) + assert not eligible(tokenizer, "{{ messages }}") + + +def test_mutated_compiled_environment_bypasses(tokenizer, monkeypatch): + from transformers.utils.chat_template_utils import _compile_jinja_template + + template = "{{ messages|length }}" + assert eligible(tokenizer, template) + env = _compile_jinja_template(template).environment + monkeypatch.setitem(env.filters, "length", lambda value: 42) + assert not eligible(tokenizer, template) + + +def test_other_stock_helper_cannot_replace_pure_helper(tokenizer, monkeypatch): + from transformers.utils.chat_template_utils import _compile_jinja_template + + template = "{{ messages|tojson }}" + assert eligible(tokenizer, template) + env = _compile_jinja_template(template).environment + monkeypatch.setitem(env.filters, "tojson", env.globals["strftime_now"]) + assert not eligible(tokenizer, template) + + +def test_generation_extension_override_bypasses(tokenizer, monkeypatch): + from transformers.utils.chat_template_utils import _compile_jinja_template + + template = "{% generation %}{{ messages }}{% endgeneration %}" + assert eligible(tokenizer, template) + env = _compile_jinja_template(template).environment + tracker = next( + ext for ext in env.extensions.values() if hasattr(ext, "_generation_support") + ) + monkeypatch.setattr( + tracker, "_generation_support", lambda *args, **kwargs: "changed" + ) + assert not eligible(tokenizer, template) From ca8e021d6a53831bab4b7524b916694c201e67f6 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 05:31:33 +0000 Subject: [PATCH 2/2] Reject customized HF dispatch and Jinja environments from render reuse --- src/art/trajectories/_render_cache.py | 15 +++- .../trajectories/test_prefix_render_cache.py | 82 +++++++++++++------ .../test_render_cache_eligibility.py | 55 ++++++++++++- 3 files changed, 121 insertions(+), 31 deletions(-) diff --git a/src/art/trajectories/_render_cache.py b/src/art/trajectories/_render_cache.py index db976f1d0..a78fc7396 100644 --- a/src/art/trajectories/_render_cache.py +++ b/src/art/trajectories/_render_cache.py @@ -48,6 +48,7 @@ def cacheable_chat_template(tokenizer, template, tools, kwargs, messages) -> boo module = sys.modules.get(cls.__module__) if ( not isinstance(tokenizer, base) + or base_module.render_jinja_template is not chat.render_jinja_template or not cls.__module__.startswith("transformers.") or getattr(module, cls.__name__, None) is not cls or type(tokenizer.chat_template) not in (str, type(None)) @@ -74,14 +75,22 @@ def cacheable_chat_template(tokenizer, template, tools, kwargs, messages) -> boo if type(kwargs) is not dict or kwargs.get("continue_final_message"): return False - from jinja2 import defaults, nodes - from jinja2.runtime import LoopContext + from jinja2 import Environment, Undefined, defaults, nodes + from jinja2.runtime import Context, LoopContext from jinja2.sandbox import ImmutableSandboxedEnvironment from jinja2.utils import Namespace compiled = chat._compile_jinja_template(template) env = compiled.environment - if type(env) is not ImmutableSandboxedEnvironment: + if ( + type(env) is not ImmutableSandboxedEnvironment + or env.undefined is not Undefined + or env.finalize is not None + or env.context_class is not Context + or env.concat is not Environment.concat + or env.autoescape is not False + or env.is_async + ): return False tree = env.parse(template) # HF's environment understands {% generation %}. parents = { diff --git a/tests/unit/trajectories/test_prefix_render_cache.py b/tests/unit/trajectories/test_prefix_render_cache.py index 63f282103..05760040a 100644 --- a/tests/unit/trajectories/test_prefix_render_cache.py +++ b/tests/unit/trajectories/test_prefix_render_cache.py @@ -2,6 +2,7 @@ from datetime import datetime, timedelta import json import string +from typing import cast from openai.types.chat import ChatCompletion import pytest @@ -18,9 +19,9 @@ def restore_retokenization_warning(monkeypatch): def test_prefix_cache_preserves_order_types_generation_and_probe_context(): calls = [] - def render(messages, *, add_generation_prompt): - calls.append(deepcopy(messages)) - return json.dumps(messages) + str(add_generation_prompt) + def render(selected_messages, *, add_generation_prompt): + calls.append(deepcopy(selected_messages)) + return json.dumps(selected_messages) + str(add_generation_prompt) messages = [{"role": "user", "content": "snow雪"}, {"a": -0.0, "b": True}] cache = tokenization._PrefixChatRenderCache(render) @@ -57,12 +58,12 @@ def test_prefix_cache_bounds_storage_and_does_not_cache_failures(monkeypatch): monkeypatch.setattr(tokenization._PrefixChatRenderCache, "_MAX_ENTRIES", 2) calls = 0 - def render(messages, *, add_generation_prompt): + def render(selected_messages, *, add_generation_prompt): nonlocal calls calls += 1 - if not messages: + if not selected_messages: raise ValueError("empty history") - return "🙂" * len(messages) + ("?" if add_generation_prompt else "") + return "🙂" * len(selected_messages) + ("?" if add_generation_prompt else "") messages = [{"content": str(i)} for i in range(20)] cache = tokenization._PrefixChatRenderCache(render) @@ -80,8 +81,12 @@ def render(messages, *, add_generation_prompt): def test_cache_settings_and_mutated_messages_invalidate_previous_prefixes(): settings = {"tool": "lookup"} - def render(messages, *, add_generation_prompt): - return settings["tool"] + json.dumps(messages) + str(add_generation_prompt) + def render(selected_messages, *, add_generation_prompt): + return ( + settings["tool"] + + json.dumps(selected_messages) + + str(add_generation_prompt) + ) messages = [{"role": "user", "content": "first"}] cache = tokenization._PrefixChatRenderCache(render) @@ -99,8 +104,8 @@ def render(messages, *, add_generation_prompt): def test_prefix_deltas_do_not_retain_quadratic_text(): - def render(messages, *, add_generation_prompt): - return "".join(m["content"] for m in messages) + ( + def render(selected_messages, *, add_generation_prompt): + return "".join(m["content"] for m in selected_messages) + ( "?" if add_generation_prompt else "" ) @@ -117,7 +122,7 @@ def render(messages, *, add_generation_prompt): @pytest.mark.parametrize("value", [object(), float("nan"), float("inf"), {1: "x"}]) def test_prefix_cache_bypasses_non_json_context(value): - def render(messages, *, add_generation_prompt): + def render(selected_messages, *, add_generation_prompt): return "unchanged" cache = tokenization._PrefixChatRenderCache(render) @@ -132,11 +137,11 @@ def test_later_generation_split_is_recomputed_when_completed_suffix_is_equal(): {"role": "assistant", "content": "HELLO"}, ] - def render(selected, *, add_generation_prompt): - text = "".join(f"<{m['role']}>{m['content']}!" for m in selected) + def render(selected_messages, *, add_generation_prompt): + text = "".join(f"<{m['role']}>{m['content']}!" for m in selected_messages) if add_generation_prompt: text += "" - if any(m["content"] == "B" for m in selected): + if any(m["content"] == "B" for m in selected_messages): text += "H" return text @@ -164,7 +169,7 @@ def render(selected, *, add_generation_prompt): _TEMPLATE = """{% for message in messages %}<{{ message.role }}>{{ message.content or '' }} -{% if message.reasoning_content %}{{ message.reasoning_content }}{% endif %} +{% if message.reasoning_content and (enable_thinking is not defined or enable_thinking) %}{{ message.reasoning_content }}{% endif %} {% for call in message.tool_calls or [] %}{% set tool_call = call.function %}{{ tool_call.name }}({% for k, v in tool_call.arguments.items() %}{{ k }}={{ v|tojson }};{% endfor %}){% endfor %} {% if message.role == 'assistant' %}${% else %}{% endif %}{% endfor %} {% if add_generation_prompt %}{% endif %}""" @@ -187,7 +192,9 @@ def _tokenizer(template): ) -def _history(turns, *, reasoning=False, refusal=False, length=False, tokenizer=None): +def _history( + turns, *, reasoning=False, refusal=False, length=False, tokenizer=None, tools=True +): messages = [{"role": "user", "content": "snow雪🙂"}] exchanges = [] for i in range(turns): @@ -206,10 +213,16 @@ def _history(turns, *, reasoning=False, refusal=False, length=False, tokenizer=N message["reasoning"] = "consider 雪" if refusal: message["refusal"] = "cannot do that" + if not tools: + del message["tool_calls"] choice = { "index": 0, "message": message, - "finish_reason": "length" if length and i == turns - 1 else "tool_calls", + "finish_reason": "length" + if length and i == turns - 1 + else "tool_calls" + if tools + else "stop", } if tokenizer is not None: @@ -256,8 +269,9 @@ def tokens(selected, generation): start = datetime(2026, 1, 1) + timedelta(seconds=i) exchanges.append( tr.ChatCompletionsExchange( - request=tr.ChatCompletionsRequest( - model="test/model", messages=deepcopy(messages) + request=cast( + tr.ChatCompletionsRequest, + {"model": "test/model", "messages": deepcopy(messages)}, ), response=response, start_time=start, @@ -265,7 +279,12 @@ def tokens(selected, generation): ) ) messages.extend( - [message, {"role": "tool", "tool_call_id": f"call-{i}", "content": "ok"}] + [ + message, + {"role": "tool", "tool_call_id": f"call-{i}", "content": "ok"} + if tools + else {"role": "user", "content": "next"}, + ] ) return tr.Trajectory( exchanges=tr.TrajectoryExchanges(chat_completions=exchanges) @@ -282,8 +301,9 @@ def tokens(selected, generation): ], ) @pytest.mark.parametrize("ends_with_assistant", [False, True]) +@pytest.mark.parametrize("enable_thinking", [False, True]) def test_cached_tool_probes_match_all_tokenized_fields( - monkeypatch, reasoning, refusal, length, ends_with_assistant + monkeypatch, reasoning, refusal, length, ends_with_assistant, enable_thinking ): history = _history(6, reasoning=reasoning, refusal=refusal, length=length) if not ends_with_assistant: @@ -298,7 +318,7 @@ def test_cached_tool_probes_match_all_tokenized_fields( tokenizer=tokenizer, base_model=None, chat_template=_TEMPLATE, - chat_template_kwargs=None, + chat_template_kwargs={"enable_thinking": enable_thinking}, _trace=trace, ) monkeypatch.setattr( @@ -312,7 +332,7 @@ def test_cached_tool_probes_match_all_tokenized_fields( tokenizer=tokenizer, base_model=None, chat_template=_TEMPLATE, - chat_template_kwargs=None, + chat_template_kwargs={"enable_thinking": enable_thinking}, _trace=expected_trace, ) assert actual.model_dump_json() == expected.model_dump_json() @@ -339,11 +359,15 @@ def test_cached_probes_preserve_native_logprobs_stop_and_source_bindings(monkeyp @pytest.mark.parametrize("turns", [4, 8, 16]) -def test_tool_probe_scaling_reuses_only_unchanged_prefixes(monkeypatch, turns): +@pytest.mark.parametrize("tools", [False, True]) +def test_probe_scaling_reuses_only_unchanged_prefixes(monkeypatch, turns, tools): from transformers import tokenization_utils_base - history = _history(turns) + history = _history(turns, tools=tools) tokenizer = _tokenizer(_TEMPLATE) + assert tokenization.cacheable_chat_template( + tokenizer, _TEMPLATE, history.tools, {}, history.messages + ) original = tokenization_utils_base.render_jinja_template calls = 0 @@ -353,6 +377,9 @@ def counted(*args, **kwargs): return original(*args, **kwargs) monkeypatch.setattr(tokenization_utils_base, "render_jinja_template", counted) + # The real gate correctly rejects custom dispatch. Admit only this test's + # pure counter after checking that the uninstrumented renderer is eligible. + monkeypatch.setattr(tokenization, "cacheable_chat_template", lambda *args: True) actual = history.tokenize(tokenizer=tokenizer, chat_template=_TEMPLATE) cached_calls = calls calls = 0 @@ -365,4 +392,7 @@ def counted(*args, **kwargs): assert actual.model_dump_json() == expected.model_dump_json() # This removes repeated unchanged prefixes, not the remaining quadratic # changed-prefix work. A zero-use cache must not pass this regression. - assert cached_calls <= calls - turns * (turns - 1) + if tools: + assert cached_calls <= calls - turns * (turns - 1) + else: + assert cached_calls == calls # Plain text already avoids tool probes. diff --git a/tests/unit/trajectories/test_render_cache_eligibility.py b/tests/unit/trajectories/test_render_cache_eligibility.py index c9f4504bd..56af42446 100644 --- a/tests/unit/trajectories/test_render_cache_eligibility.py +++ b/tests/unit/trajectories/test_render_cache_eligibility.py @@ -34,7 +34,7 @@ def test_stock_render_with_local_macro_namespace_and_generation(tokenizer): {% generation %}{{ content(message) }}{% endgeneration %}{% endfor %} {% for key, value in tools[0]|items %}{{ key }}={{ value|tojson }}{% endfor %} {{ ns.n }}{% if add_generation_prompt %}assistant{% endif %}""" - tools = [{"name": "lookup"}] + tools: list[dict[str, object]] = [{"name": "lookup"}] assert eligible(tokenizer, template, tools=tools) assert "hello" in tokenizer.apply_chat_template( [{"role": "user", "content": "hello"}], @@ -78,7 +78,7 @@ def test_non_plain_context_bypasses(tokenizer, field, value): def test_mutable_context_is_rechecked_and_cycles_bypass(tokenizer): - tools = [{"name": "lookup"}] + tools: list[dict[str, object]] = [{"name": "lookup"}] assert eligible(tokenizer, "{{ tools|tojson }}", tools=tools) tools[0]["callback"] = lambda: None assert not eligible(tokenizer, "{{ tools|tojson }}", tools=tools) @@ -125,6 +125,57 @@ def test_mutated_compiled_environment_bypasses(tokenizer, monkeypatch): assert not eligible(tokenizer, template) +def test_custom_hf_render_dispatch_bypasses_without_invoking_it(tokenizer, monkeypatch): + from transformers import tokenization_utils_base + + assert eligible(tokenizer, "{{ messages }}") + + def custom(*args, **kwargs): + raise AssertionError("eligibility must not invoke a custom dispatch") + + monkeypatch.setattr(tokenization_utils_base, "render_jinja_template", custom) + assert not eligible(tokenizer, "{{ messages }}") + + +def test_custom_undefined_changes_render_and_bypasses(tokenizer, monkeypatch): + from jinja2 import Undefined + from transformers.utils.chat_template_utils import _compile_jinja_template + + class CountingUndefined(Undefined): + calls = 0 + + def __str__(self): + type(self).calls += 1 + return str(self.calls) + + template = "{{ messages[0].missing }}" + assert eligible(tokenizer, template) + env = _compile_jinja_template(template).environment + monkeypatch.setattr(env, "undefined", CountingUndefined) + assert env.from_string(template).render(messages=[{}]) == "1" + assert _compile_jinja_template(template).render(messages=[{}]) == "2" + assert not eligible(tokenizer, template) + + +@pytest.mark.parametrize( + "name,value", + [ + ("finalize", lambda value: value), + ("context_class", object), + ("concat", lambda values: "".join(values)), + ("autoescape", True), + ("is_async", True), + ], +) +def test_custom_environment_configuration_bypasses(tokenizer, monkeypatch, name, value): + from transformers.utils.chat_template_utils import _compile_jinja_template + + template = "{{ messages }}" + assert eligible(tokenizer, template) + monkeypatch.setattr(_compile_jinja_template(template).environment, name, value) + assert not eligible(tokenizer, template) + + def test_other_stock_helper_cannot_replace_pure_helper(tokenizer, monkeypatch): from transformers.utils.chat_template_utils import _compile_jinja_template