diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py
index 61626c6fc..a51e2b436 100644
--- a/src/art/trajectories/_tokenize.py
+++ b/src/art/trajectories/_tokenize.py
@@ -113,7 +113,12 @@ def __init__(self, render: _ChatRender) -> None:
self.bytes = 0
def for_messages(
- self, messages: list[dict[str, Any]], text: str, *, settings: object = None
+ self,
+ messages: list[dict[str, Any]],
+ text: str,
+ *,
+ settings: object = None,
+ add_generation_prompt: bool | None = None,
) -> _ChatRender:
try:
context = tuple(_render_context_key(message) for message in messages)
@@ -130,10 +135,23 @@ def for_messages(
break
common += 1
+ rendered_generation_prompt = add_generation_prompt
+
def render(
selected_messages: list[dict[str, Any]], *, add_generation_prompt: bool
) -> str:
count = len(selected_messages)
+ # This exact full-context render was just evaluated by the caller.
+ # A changed probe still proves every other generation/completion.
+ if (
+ count == len(messages)
+ and add_generation_prompt == rendered_generation_prompt
+ and all(
+ original is current
+ for original, current in zip(messages, selected_messages)
+ )
+ ):
+ return text
if count > common or any(
original is not current
for original, current in zip(messages, selected_messages)
@@ -5277,6 +5295,7 @@ def segmented_render(
span_render = prefix_render_cache.for_messages(
selected_messages,
text,
+ add_generation_prompt=add_generation_prompt,
settings=_render_context_key(
[
history.tools,
diff --git a/tests/unit/trajectories/test_prefix_render_cache.py b/tests/unit/trajectories/test_prefix_render_cache.py
index 05760040a..a285edb96 100644
--- a/tests/unit/trajectories/test_prefix_render_cache.py
+++ b/tests/unit/trajectories/test_prefix_render_cache.py
@@ -129,7 +129,10 @@ def render(selected_messages, *, add_generation_prompt):
assert cache.for_messages([{"content": value}], "unchanged") is render
-def test_later_generation_split_is_recomputed_when_completed_suffix_is_equal():
+@pytest.mark.parametrize("seed_completed", [False, True])
+def test_later_generation_split_is_recomputed_when_completed_suffix_is_equal(
+ seed_completed,
+):
messages = [
{"role": "user", "content": "x"},
{"role": "assistant", "content": "A"},
@@ -150,7 +153,11 @@ def render(selected_messages, *, add_generation_prompt):
tokenization._assistant_char_spans(
messages,
original,
- cache.for_messages(messages, original),
+ cache.for_messages(
+ messages,
+ original,
+ add_generation_prompt=False if seed_completed else None,
+ ),
add_generation_prompt=False,
)
probe = deepcopy(messages)
@@ -158,7 +165,12 @@ def render(selected_messages, *, add_generation_prompt):
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
+ probe,
+ text,
+ cache.for_messages(
+ probe, text, add_generation_prompt=False if seed_completed else None
+ ),
+ add_generation_prompt=False,
)
expected = tokenization._assistant_char_spans(
probe, text, render, add_generation_prompt=False
@@ -168,6 +180,29 @@ def render(selected_messages, *, add_generation_prompt):
assert text[start:end] == "ELLO!"
+def test_completed_render_reuse_requires_same_messages_and_generation():
+ calls = []
+
+ def render(selected_messages, *, add_generation_prompt):
+ calls.append((deepcopy(selected_messages), add_generation_prompt))
+ return json.dumps(selected_messages) + str(add_generation_prompt)
+
+ messages = [{"role": "assistant", "content": "original"}]
+ cache = tokenization._PrefixChatRenderCache(render)
+ cache.for_messages(messages, render(messages, add_generation_prompt=False))
+ probe = [{"role": "assistant", "content": "changed"}]
+ text = render(probe, add_generation_prompt=False)
+ cached = cache.for_messages(probe, text, add_generation_prompt=False)
+ before = len(calls)
+ assert cached(probe[:], add_generation_prompt=False) == text
+ assert len(calls) == before
+ assert cached(probe, add_generation_prompt=True) != text
+ assert len(calls) == before + 1
+ assert cached(deepcopy(probe), add_generation_prompt=False) == text
+ assert len(calls) == before + 2
+ assert cached([], add_generation_prompt=False) == "[]False"
+
+
_TEMPLATE = """{% for message in messages %}<{{ message.role }}>{{ message.content or '' }}
{% 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 %}
@@ -395,4 +430,4 @@ def counted(*args, **kwargs):
if tools:
assert cached_calls <= calls - turns * (turns - 1)
else:
- assert cached_calls == calls # Plain text already avoids tool probes.
+ assert cached_calls == calls - 1 # The completed full render is reused.