diff --git a/dev/trainer_rank.py b/dev/trainer_rank.py
index 0bfc45eb8..65aa2c6ff 100644
--- a/dev/trainer_rank.py
+++ b/dev/trainer_rank.py
@@ -5,9 +5,9 @@
import torch
import torch.distributed as dist
from trainer_rank_support import load_random_checkpoints
-from transformers import AutoTokenizer
import typer
+from art import get_tokenizer
from art.trainer_rank import AdamParams, ForwardInput, TrainerRank
@@ -34,9 +34,7 @@ def main(
from art.megatron import train as megatron_train
- tokenizer = cast(
- Any, AutoTokenizer.from_pretrained(model, trust_remote_code=True)
- )
+ tokenizer = cast(Any, get_tokenizer(model, trust_remote_code=True))
inputs: list[ForwardInput[torch.Tensor, None, None, None]] = []
rows = load_dataset("roneneldan/TinyStories", split="train", streaming=True)
for row in islice(rows, samples):
diff --git a/dev/trainer_rank_landing_acceptance.py b/dev/trainer_rank_landing_acceptance.py
index 95513fd1d..dadd80c97 100644
--- a/dev/trainer_rank_landing_acceptance.py
+++ b/dev/trainer_rank_landing_acceptance.py
@@ -197,13 +197,11 @@ def _check_corpus_tokenizer(corpus: dict[str, Any], model: str) -> dict[str, Any
compares vocabulary sizes; it fails the cell loudly on any difference.
"""
- from transformers import AutoTokenizer
+ from art import get_tokenizer
corpus_model = str(corpus.get("tokenizer_model") or "")
- model_tokenizer = AutoTokenizer.from_pretrained(model, trust_remote_code=True)
- corpus_tokenizer = AutoTokenizer.from_pretrained(
- corpus_model, trust_remote_code=True
- )
+ model_tokenizer = get_tokenizer(model, trust_remote_code=True)
+ corpus_tokenizer = get_tokenizer(corpus_model, trust_remote_code=True)
sample = corpus["groups"][0]["histories"][0]["tokens"][:256]
problems = []
if len(model_tokenizer) != len(corpus_tokenizer):
diff --git a/examples/hn_title_generator/train.py b/examples/hn_title_generator/train.py
index f2f0b98d4..2759834fa 100644
--- a/examples/hn_title_generator/train.py
+++ b/examples/hn_title_generator/train.py
@@ -8,10 +8,10 @@
import openai
from openai.types.chat import ChatCompletionMessageParam
from openpipe import AsyncOpenPipe
-from transformers.models.auto.tokenization_auto import AutoTokenizer
from utils import cache, prompt_for_title, pull_data, score_title
import art
+from art import get_tokenizer
from art.local import LocalBackend
from art.utils import iterate_dataset, limit_concurrency
@@ -37,7 +37,7 @@ def filter_on_length(data: Dataset, max_length: int, tokenizer_name: str) -> Dat
print(
f"Filtering dataset for max prompt length: {max_length} using tokenizer: {tokenizer_name}"
)
- tokenizer = AutoTokenizer.from_pretrained(tokenizer_name)
+ tokenizer = get_tokenizer(tokenizer_name)
def check_length(x):
# Ensure 'prompt' is a list of dicts
diff --git a/pyproject.toml b/pyproject.toml
index 96e66f7e6..82a527972 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -152,7 +152,7 @@ build-backend = "hatchling.build"
allow-direct-references = true
[tool.hatch.build.targets.wheel]
-packages = ["src/art", "src/mp_actors"]
+packages = ["src/art", "src/art_inference", "src/mp_actors"]
[tool.hatch.build.targets.wheel.force-include]
".agents/skills" = "art/skills"
diff --git a/src/art/__init__.py b/src/art/__init__.py
index a5fdc3f19..fb8fe9638 100644
--- a/src/art/__init__.py
+++ b/src/art/__init__.py
@@ -71,6 +71,7 @@
from .model import Model, TrainableModel
from .pipeline_tuner import PipelineAutotuneConfig, PipelineRuntimeConfig
from .serverless import ServerlessBackend
+from .tokenizer import get_tokenizer
from .trajectories import (
Trajectory,
TrajectoryGroup,
@@ -116,6 +117,7 @@
"PIPELINE_RL_METRIC_DEFINITIONS",
"PIPELINE_RL_SCORE_METRICS",
"get_megatron_runtime_config",
+ "get_tokenizer",
"init_megatron_runtime_config",
"ServerlessBackend",
"ServerlessTrainResult",
diff --git a/src/art/local/backend.py b/src/art/local/backend.py
index 51cb18358..cc5f72b8e 100644
--- a/src/art/local/backend.py
+++ b/src/art/local/backend.py
@@ -16,9 +16,9 @@
from typing import TYPE_CHECKING, Any, AsyncIterator, Iterable, Literal, cast
import warnings
+from art.tokenizer import get_tokenizer
from art.utils.chat_template import (
chat_template_with_preserved_thinking,
- configure_preserved_thinking_chat_template,
)
from art.utils.lifecycle import (
PROCESS_SHUTDOWN_TIMEOUT_SECONDS,
@@ -40,7 +40,6 @@
from pydantic import BaseModel, ConfigDict
import torch
from tqdm import auto as tqdm
-from transformers import AutoTokenizer
from transformers.tokenization_utils_base import PreTrainedTokenizerBase
from typing_extensions import Self
@@ -196,7 +195,9 @@ def _apply_configured_chat_template(
) -> None:
chat_template = _configured_chat_template_value(internal_config)
if chat_template is not None:
- tokenizer.chat_template = chat_template
+ tokenizer.chat_template = cast(
+ str, chat_template_with_preserved_thinking(chat_template)
+ )
def _model_support_handler(
@@ -241,9 +242,9 @@ def _apply_configured_chat_template_server_args(
chat_template = _model_support_default_chat_template(
base_model, internal_config
)
- if chat_template is None and _should_probe_preserve_thinking_template(base_model):
+ if chat_template is None and base_model is not None:
try:
- tokenizer = AutoTokenizer.from_pretrained(base_model)
+ tokenizer = get_tokenizer(base_model)
except (OSError, ValueError) as error:
warnings.warn(
f"Could not load {base_model!r} to configure prior-thinking "
@@ -252,13 +253,14 @@ def _apply_configured_chat_template_server_args(
stacklevel=2,
)
else:
- default = getattr(tokenizer, "chat_template", None)
- preserved = chat_template_with_preserved_thinking(default)
- if preserved != default:
- chat_template = cast(str, preserved)
+ template = getattr(tokenizer, "chat_template", None)
+ if isinstance(template, str) and ("{{" in template or "{%" in template):
+ chat_template = template
if chat_template is None:
return
- server_args.setdefault("chat_template", chat_template)
+ server_args.setdefault(
+ "chat_template", chat_template_with_preserved_thinking(chat_template)
+ )
if chat_template_content_format := internal_config.get(
"chat_template_content_format"
):
@@ -269,13 +271,6 @@ def _apply_configured_chat_template_server_args(
config_dict["server_args"] = server_args
-def _should_probe_preserve_thinking_template(base_model: str | None) -> bool:
- if base_model is None:
- return False
- model_name = base_model.rstrip("/").rsplit("/", 1)[-1]
- return model_name.startswith(("Qwen3-", "Qwen3.5-"))
-
-
def _tokenizer_cache_key(
base_model: str,
internal_config: dev.InternalModelConfig,
@@ -287,12 +282,7 @@ def _tokenizer_cache_key(
def _load_training_tokenizer(base_model: str) -> PreTrainedTokenizerBase:
- return cast(
- PreTrainedTokenizerBase,
- configure_preserved_thinking_chat_template(
- AutoTokenizer.from_pretrained(base_model)
- ),
- )
+ return get_tokenizer(base_model)
class LocalBackend:
diff --git a/src/art/megatron/dsv4/tokenizer.py b/src/art/megatron/dsv4/tokenizer.py
index 03d44b643..1c59a19f5 100644
--- a/src/art/megatron/dsv4/tokenizer.py
+++ b/src/art/megatron/dsv4/tokenizer.py
@@ -1,11 +1,18 @@
from __future__ import annotations
+from contextvars import ContextVar
import copy
from typing import Any
from transformers.tokenization_utils_base import PreTrainedTokenizerBase
-from art.megatron.dsv4.encoding import encode_messages
+from art.megatron.dsv4 import encoding
+from art_inference.append_only import patch_deepseek_renderer, preserves_history
+
+_PRESERVE_HISTORY: ContextVar[bool] = ContextVar(
+ "art_dsv4_preserve_history", default=True
+)
+patch_deepseek_renderer(encoding, _PRESERVE_HISTORY.get)
DSV4_CHAT_TEMPLATE_MARKER = "deepseek_v4_python_encoder enable_thinking"
@@ -35,12 +42,28 @@ def apply_chat_template(
tools: list[dict[str, Any]] | None = None,
**kwargs: Any,
) -> str | list[int]:
+ chat_template = kwargs.get("chat_template")
+ if chat_template is None:
+ chat_template = self.chat_template
+ if chat_template != DSV4_CHAT_TEMPLATE_MARKER:
+ return super().apply_chat_template(messages, tools=tools, **kwargs)
thinking = bool(kwargs.get("thinking", False)) or bool(
kwargs.get("enable_thinking", False)
)
thinking_mode = "thinking" if thinking else "chat"
conversation = kwargs.get("conversation", messages)
- rendered_messages = list(conversation)
+ rendered_messages = [
+ {
+ **message,
+ "reasoning": message.get(
+ "reasoning",
+ message.get("reasoning_content", message.get("thinking")),
+ ),
+ }
+ if message.get("role") == "assistant"
+ else message
+ for message in conversation
+ ]
if tools:
rendered_messages.insert(0, {"role": "system", "tools": tools})
@@ -55,12 +78,17 @@ def apply_chat_template(
else:
reasoning_effort = "high"
- prompt = encode_messages(
- rendered_messages,
- thinking_mode=thinking_mode,
- drop_thinking=kwargs.get("drop_thinking", True),
- reasoning_effort=reasoning_effort,
- )
+ preserve = preserves_history(kwargs)
+ token = _PRESERVE_HISTORY.set(preserve)
+ try:
+ prompt = encoding.encode_messages(
+ rendered_messages,
+ thinking_mode=thinking_mode,
+ drop_thinking=not preserve,
+ reasoning_effort=reasoning_effort,
+ )
+ finally:
+ _PRESERVE_HISTORY.reset(token)
if not kwargs.get("tokenize", True):
return prompt
tokenizer_kwargs = {
diff --git a/src/art/tinker/prefix_cache.py b/src/art/tinker/prefix_cache.py
deleted file mode 100644
index 81c8f4620..000000000
--- a/src/art/tinker/prefix_cache.py
+++ /dev/null
@@ -1,208 +0,0 @@
-from __future__ import annotations
-
-from collections import OrderedDict
-from dataclasses import dataclass
-from typing import Sequence
-
-
-@dataclass(frozen=True)
-class PrefixEntry:
- rendered_len: int
- raw_prefix: tuple[int, ...]
-
-
-@dataclass
-class PrefixCacheStats:
- max_entries: int
- lookups: int = 0
- hits: int = 0
- misses: int = 0
- inserts: int = 0
- replaced_entries: int = 0
- evictions: int = 0
- splits: int = 0
- pruned_nodes: int = 0
- merged_nodes: int = 0
- lru_repairs: int = 0
-
-
-class _RadixEdge:
- __slots__ = ("label", "child")
-
- def __init__(self, label: tuple[int, ...], child: _RadixNode) -> None:
- self.label = label
- self.child = child
-
-
-class _RadixNode:
- __slots__ = ("entry", "children", "parent", "parent_token")
-
- def __init__(
- self, parent: _RadixNode | None = None, parent_token: int | None = None
- ) -> None:
- self.entry: PrefixEntry | None = None
- self.children: dict[int, _RadixEdge] = {}
- self.parent = parent
- self.parent_token = parent_token
-
-
-def _common_prefix_len(
- tokens: Sequence[int], start: int, label: tuple[int, ...]
-) -> int:
- max_len = min(len(tokens) - start, len(label))
- i = 0
- while i < max_len and tokens[start + i] == label[i]:
- i += 1
- return i
-
-
-class LRUTrieCache:
- """LRU-bounded radix trie for token sequence rewrites."""
-
- def __init__(self, max_entries: int = 16_384) -> None:
- if max_entries <= 0:
- raise ValueError("max_entries must be positive")
- self._root = _RadixNode()
- self._lru: OrderedDict[_RadixNode, None] = OrderedDict()
- self._max_entries = max_entries
- self.stats = PrefixCacheStats(max_entries=max_entries)
-
- def lookup(self, rendered_tokens: Sequence[int]) -> PrefixEntry | None:
- self.stats.lookups += 1
- node = self._root
- idx = 0
- best_node = None
- while idx < len(rendered_tokens):
- edge = node.children.get(rendered_tokens[idx])
- if edge is None:
- break
- matched = _common_prefix_len(rendered_tokens, idx, edge.label)
- if matched != len(edge.label):
- break
- idx += matched
- node = edge.child
- if node.entry is not None:
- best_node = node
- if best_node is None:
- self.stats.misses += 1
- return None
- self.stats.hits += 1
- try:
- self._lru.move_to_end(best_node)
- except KeyError:
- self.stats.lru_repairs += 1
- self._lru[best_node] = None
- self._lru.move_to_end(best_node)
- self._evict()
- return best_node.entry
-
- def insert(self, rendered_prefix: Sequence[int], raw_prefix: Sequence[int]) -> None:
- self.stats.inserts += 1
- node = self._root
- idx = 0
- while idx < len(rendered_prefix):
- token = rendered_prefix[idx]
- edge = node.children.get(token)
- if edge is None:
- child = _RadixNode(parent=node, parent_token=token)
- node.children[token] = _RadixEdge(tuple(rendered_prefix[idx:]), child)
- node = child
- idx = len(rendered_prefix)
- break
-
- matched = _common_prefix_len(rendered_prefix, idx, edge.label)
- if matched == len(edge.label):
- idx += matched
- node = edge.child
- continue
-
- mid = _RadixNode(parent=node, parent_token=token)
- self.stats.splits += 1
- old_suffix = edge.label[matched:]
- old_child = edge.child
- old_child.parent = mid
- old_child.parent_token = old_suffix[0]
- mid.children[old_suffix[0]] = _RadixEdge(old_suffix, old_child)
- edge.label = edge.label[:matched]
- edge.child = mid
- node = mid
- idx += matched
- if idx < len(rendered_prefix):
- new_token = rendered_prefix[idx]
- child = _RadixNode(parent=node, parent_token=new_token)
- node.children[new_token] = _RadixEdge(
- tuple(rendered_prefix[idx:]), child
- )
- node = child
- break
-
- if node.entry is not None:
- self.stats.replaced_entries += 1
- node.entry = PrefixEntry(
- rendered_len=len(rendered_prefix), raw_prefix=tuple(raw_prefix)
- )
- self._lru[node] = None
- self._lru.move_to_end(node)
- self._evict()
-
- def _evict(self) -> None:
- while len(self._lru) > self._max_entries:
- old_node, _ = self._lru.popitem(last=False)
- self.stats.evictions += 1
- old_node.entry = None
- self._prune(old_node)
-
- def _prune(self, node: _RadixNode) -> None:
- # Collapse empty branches after eviction so the bounded cache stays bounded.
- while node.parent is not None:
- parent = node.parent
- parent_token = node.parent_token
- assert parent_token is not None
-
- if node.entry is None and not node.children:
- del parent.children[parent_token]
- self.stats.pruned_nodes += 1
- node = parent
- continue
-
- if node.entry is None and len(node.children) == 1:
- _, child_edge = next(iter(node.children.items()))
- parent_edge = parent.children[parent_token]
- parent_edge.label = parent_edge.label + child_edge.label
- parent_edge.child = child_edge.child
- child_edge.child.parent = parent
- child_edge.child.parent_token = parent_token
- self.stats.merged_nodes += 1
- node = parent
- continue
-
- break
-
- def snapshot_stats(self) -> dict[str, int | float]:
- hit_rate = self.stats.hits / self.stats.lookups if self.stats.lookups else 0.0
- return {
- "enabled": True,
- "max_entries": self.stats.max_entries,
- "current_entries": len(self._lru),
- "node_count": self._node_count(),
- "lookups": self.stats.lookups,
- "hits": self.stats.hits,
- "misses": self.stats.misses,
- "hit_rate": hit_rate,
- "inserts": self.stats.inserts,
- "replaced_entries": self.stats.replaced_entries,
- "evictions": self.stats.evictions,
- "splits": self.stats.splits,
- "pruned_nodes": self.stats.pruned_nodes,
- "merged_nodes": self.stats.merged_nodes,
- "lru_repairs": self.stats.lru_repairs,
- }
-
- def _node_count(self) -> int:
- count = 0
- stack = [self._root]
- while stack:
- node = stack.pop()
- count += 1
- stack.extend(edge.child for edge in node.children.values())
- return count
diff --git a/src/art/tinker/server.py b/src/art/tinker/server.py
index 1767fa868..0909d58b9 100644
--- a/src/art/tinker/server.py
+++ b/src/art/tinker/server.py
@@ -29,14 +29,24 @@
from pydantic import BaseModel, Field, SkipValidation, TypeAdapter
import tinker
from tinker_cookbook import renderers
-from tinker_cookbook.tokenizer_utils import get_tokenizer
+from tinker_cookbook.tokenizer_utils import Tokenizer as CookbookTokenizer
from transformers.tokenization_utils_base import BatchEncoding
import uvicorn
-from art.tinker.prefix_cache import LRUTrieCache
from art.tinker.renderers import get_renderer_name, is_qwen3_dot_family_model
+from art.token_prefix import TokenPrefixStore
+from art.tokenizer import get_tokenizer
from art.types import Message, Tools
-from art.utils.chat_template import default_chat_template_kwargs_for_tokenizer
+from art.utils.append_only import (
+ chat_prefix_observations,
+ has_renderable_tool_arguments,
+ output_prefix_observations,
+ preserves_history,
+)
+from art.utils.chat_template import (
+ default_chat_template_kwargs_for_tokenizer,
+ normalize_tool_call_arguments_for_chat_template,
+)
from mp_actors import close_proxy, move_to_child_process
@@ -120,7 +130,7 @@ class OpenAICompatibleTinkerServer:
port: int | None = None
num_workers: int | None = None
max_concurrent_sampling_clients: int | None = None
- _prefix_cache: LRUTrieCache = field(default_factory=LRUTrieCache)
+ _prefix_cache: TokenPrefixStore = field(default_factory=TokenPrefixStore)
_task: asyncio.Task[None] | None = None
_tenants: dict[str, "OpenAICompatibleTinkerServerTenant"] = field(
default_factory=dict
@@ -366,19 +376,30 @@ async def chat_completions(
worker = next(workers)
tenant = self._get_request_tenant(request)
samplable_model = await tenant.get_samplable_model(body["model"])
+ template_kwargs = cast(dict[str, Any], body).get("chat_template_kwargs")
+ preserve = preserves_history(template_kwargs)
+ scope = json.dumps(
+ [id(tenant), samplable_model.base_model, template_kwargs],
+ sort_keys=True,
+ )
rendered_prompt_tokens = await worker.prompt_tokens(
base_model=samplable_model.base_model,
messages=list(body["messages"]),
tools=list(body.get("tools", [])) if "tools" in body else None,
+ chat_template_kwargs=template_kwargs,
)
prompt_tokens = rendered_prompt_tokens
- prefix_entry = self._prefix_cache.lookup(rendered_prompt_tokens)
- if prefix_entry is not None and prefix_entry.rendered_len <= len(
+ prefix_entry = (
+ self._prefix_cache.lookup(scope, rendered_prompt_tokens, None)
+ if preserve
+ else None
+ )
+ if prefix_entry is not None and prefix_entry.rendered_length <= len(
rendered_prompt_tokens
):
prompt_tokens = (
list(prefix_entry.raw_prefix)
- + rendered_prompt_tokens[prefix_entry.rendered_len :]
+ + rendered_prompt_tokens[prefix_entry.rendered_length :]
)
try:
async with samplable_model.sampling_client() as sampling_client:
@@ -407,18 +428,20 @@ async def chat_completions(
raise HTTPException(status_code=e.status_code, detail=detail) from e
(
chat_completion,
- token_discrepancies,
- ) = await worker.chat_completion_and_token_discrepancies(
+ prefixes,
+ ) = await worker.chat_completion_and_prefixes(
base_model=samplable_model.base_model,
sample_response=sample_response,
model_name=body["model"],
- prompt_tokens=len(prompt_tokens),
+ prompt_tokens=prompt_tokens,
+ rendered_prompt=rendered_prompt_tokens,
+ messages=list(body["messages"]),
+ tools=list(body.get("tools", [])) if "tools" in body else None,
+ chat_template_kwargs=template_kwargs,
)
- for rendered_response_tokens, raw_response_tokens in token_discrepancies:
- self._prefix_cache.insert(
- rendered_prompt_tokens + rendered_response_tokens,
- prompt_tokens + raw_response_tokens,
- )
+ if preserve:
+ for rendered, raw, edits in prefixes:
+ self._prefix_cache.insert(scope, rendered, raw, "content", edits)
return chat_completion
server_config = uvicorn.Config(
@@ -550,14 +573,24 @@ async def prompt_tokens(
base_model: str,
messages: list[ChatCompletionMessageParam],
tools: list[ChatCompletionToolUnionParam] | None,
+ *,
+ add_generation_prompt: bool = True,
+ chat_template_kwargs: dict[str, Any] | None = None,
) -> list[int]:
normalized_messages = _normalize_qwen3_dot_messages(base_model, messages)
tokenizer = self._get_renderer(base_model).tokenizer
- chat_template_kwargs = default_chat_template_kwargs_for_tokenizer(tokenizer)
+ normalized_messages = normalize_tool_call_arguments_for_chat_template(
+ normalized_messages, getattr(tokenizer, "chat_template", None)
+ )
+ chat_template_kwargs = {
+ **default_chat_template_kwargs_for_tokenizer(tokenizer),
+ **(chat_template_kwargs or {}),
+ }
encoding = tokenizer.apply_chat_template(
cast(Any, normalized_messages),
tools=cast(Any, tools),
- add_generation_prompt=True,
+ tokenize=True,
+ add_generation_prompt=add_generation_prompt,
**chat_template_kwargs,
)
if isinstance(encoding, BatchEncoding):
@@ -590,28 +623,69 @@ async def messages_and_choices_prompt_tokens_and_choice_offsets(
)
return (result.token_ids, result.choice_offsets) if result is not None else None
- async def chat_completion_and_token_discrepancies(
+ async def chat_completion_and_prefixes(
self,
base_model: str,
sample_response: tinker.SampleResponse,
model_name: str,
- prompt_tokens: int,
- ) -> tuple[ChatCompletion, list[tuple[list[int], list[int]]]]:
+ prompt_tokens: list[int],
+ rendered_prompt: list[int],
+ messages: list[ChatCompletionMessageParam],
+ tools: list[ChatCompletionToolUnionParam] | None,
+ chat_template_kwargs: dict[str, Any] | None = None,
+ ) -> tuple[ChatCompletion, list[tuple[list[int], list[int], tuple[Any, ...]]]]:
renderer = self._get_renderer(base_model)
choices: list[Choice] = []
- token_discrepancies: list[tuple[list[int], list[int]]] = []
+ prefixes = []
for i, sequence in enumerate(sample_response.sequences):
assert sequence.logprobs is not None, "Logprobs are required"
assert len(sequence.tokens) == len(sequence.logprobs), (
"Tokens and logprobs must have the same length"
)
- rendered_response_tokens = renderer.tokenizer.encode(
- renderer.tokenizer.decode(sequence.tokens)
- )
- if rendered_response_tokens != sequence.tokens:
- token_discrepancies.append((rendered_response_tokens, sequence.tokens))
message, _ = renderer.parse_response(sequence.tokens)
openai_message = renderer.to_openai_message(message)
+ if preserves_history(chat_template_kwargs):
+ prefixes.extend(
+ output_prefix_observations(
+ renderer.tokenizer,
+ rendered_prompt,
+ prompt_tokens,
+ sequence.tokens,
+ )
+ )
+ if preserves_history(
+ chat_template_kwargs
+ ) and has_renderable_tool_arguments(openai_message):
+
+ async def render(assistant: dict[str, Any]) -> list[int]:
+ return await self.prompt_tokens(
+ base_model,
+ [*messages, cast(Any, assistant)],
+ tools,
+ add_generation_prompt=False,
+ chat_template_kwargs=chat_template_kwargs,
+ )
+
+ reasoning = {
+ key: openai_message[key]
+ for key in ("reasoning", "reasoning_content", "thinking")
+ if openai_message.get(key) is not None
+ }
+ prefixes.extend(
+ chat_prefix_observations(
+ renderer.tokenizer,
+ rendered_prompt,
+ await render(openai_message),
+ prompt_tokens,
+ sequence.tokens,
+ reasoning_prompt=await render(
+ {"role": "assistant", "content": "", **reasoning}
+ )
+ if reasoning
+ else None,
+ complete=sequence.stop_reason == "stop",
+ )
+ )
tool_calls = (
[
ChatCompletionMessageFunctionToolCall(
@@ -635,10 +709,12 @@ async def chat_completion_and_token_discrepancies(
Choice(
finish_reason=sequence.stop_reason,
index=i,
- message=ChatCompletionMessage(
- content=openai_message.get("content") or None,
- role="assistant",
- tool_calls=tool_calls, # type: ignore
+ message=ChatCompletionMessage.model_validate(
+ {
+ **openai_message,
+ "role": "assistant",
+ "tool_calls": tool_calls,
+ }
),
logprobs=ChoiceLogprobs(
content=[
@@ -669,18 +745,19 @@ async def chat_completion_and_token_discrepancies(
object="chat.completion",
usage=CompletionUsage(
completion_tokens=completion_tokens,
- prompt_tokens=prompt_tokens,
- total_tokens=completion_tokens + prompt_tokens,
+ prompt_tokens=len(prompt_tokens),
+ total_tokens=completion_tokens + len(prompt_tokens),
),
),
- token_discrepancies,
+ prefixes,
)
def _get_renderer(self, base_model: str) -> renderers.Renderer:
if base_model not in self._renderers:
self._renderers[base_model] = renderers.get_renderer(
name=get_renderer_name(base_model),
- tokenizer=get_tokenizer(base_model),
+ # Cookbook's annotation omits the fast HF tokenizers it accepts.
+ tokenizer=cast(CookbookTokenizer, get_tokenizer(base_model)),
model_name=base_model,
)
return self._renderers[base_model]
diff --git a/src/art/tinker_native/backend.py b/src/art/tinker_native/backend.py
index 8641c3318..879c113ac 100644
--- a/src/art/tinker_native/backend.py
+++ b/src/art/tinker_native/backend.py
@@ -25,10 +25,13 @@
from openai.types.chat.completion_create_params import CompletionCreateParams
from openai.types.completion_usage import CompletionUsage
import tinker
-from tinker_cookbook import renderers, tokenizer_utils
+from tinker_cookbook import renderers
+from tinker_cookbook.tokenizer_utils import Tokenizer as CookbookTokenizer
import torch
import uvicorn
+from art.tokenizer import get_tokenizer
+
from .. import dev
from ..adapter_leases import pin_inference_step, pinned_inference_step
from ..backend import Backend
@@ -715,12 +718,15 @@ async def _build_model_state(self, model: TrainableModel) -> ModelState:
service_client = tinker.ServiceClient()
rest_client = service_client.create_rest_client()
- tokenizer = tokenizer_utils.get_tokenizer(model.base_model)
+ tokenizer = get_tokenizer(model.base_model)
renderer = renderers.get_renderer(
name=config.renderer_name,
- tokenizer=tokenizer,
+ # Cookbook's annotation omits the fast HF tokenizers it accepts.
+ tokenizer=cast(CookbookTokenizer, tokenizer),
model_name=model.base_model,
)
+ if hasattr(renderer, "strip_thinking_from_history"):
+ setattr(renderer, "strip_thinking_from_history", False)
saved_state = model.read_state() or {}
tinker_run_ids = list(saved_state.get(STATE_KEY_RUN_IDS, []))
diff --git a/src/art/token_prefix.py b/src/art/token_prefix.py
new file mode 100644
index 000000000..0af1a4ce0
--- /dev/null
+++ b/src/art/token_prefix.py
@@ -0,0 +1,8 @@
+"""Compatibility imports for ART's lightweight inference helpers."""
+
+import sys
+
+from art_inference import token_prefix as _implementation
+from art_inference.token_prefix import * # noqa: F403
+
+sys.modules[__name__] = _implementation
diff --git a/src/art/tokenizer.py b/src/art/tokenizer.py
new file mode 100644
index 000000000..f3e26414e
--- /dev/null
+++ b/src/art/tokenizer.py
@@ -0,0 +1,128 @@
+"""The default tokenizer used by ART's renderers and training backends."""
+
+from __future__ import annotations
+
+from copy import copy
+import os
+import sys
+from typing import TYPE_CHECKING, Any, cast
+
+from .utils.chat_template import configure_preserved_thinking_chat_template
+
+if TYPE_CHECKING:
+ from transformers import PreTrainedTokenizerBase
+
+_KIMI_TOKENIZER_REVISIONS = {
+ "moonshotai/Kimi-K2-Thinking": "a51ccc050d73dab088bf7b0e2dd9b30ae85a4e55",
+ "moonshotai/Kimi-K2.5": "2426b45b6af0da48d0dcce71bbce6225e5c73adc",
+ "moonshotai/Kimi-K2.6": "b5aabbfb20227ed42becbf5541dbffd213942c58",
+}
+_LLAMA_TEXT_MODELS = {
+ "meta-llama/Llama-3.1-8B-Instruct",
+ "meta-llama/Llama-3.1-8B",
+ "meta-llama/Llama-3.1-70B",
+ "meta-llama/Llama-3.2-1B",
+ "meta-llama/Llama-3.2-3B",
+ "meta-llama/Llama-3.3-70B-Instruct",
+}
+_LLAMA_CHAT_TOKENIZER = "thinkingmachineslabinc/meta-llama-3-instruct-tokenizer"
+
+
+def get_tokenizer(
+ base_model: str, *, revision: str | None = None, **kwargs: Any
+) -> PreTrainedTokenizerBase:
+ """Load a model's tokenizer with ART's history-preserving defaults.
+
+ ``revision`` and other keyword arguments go to ``from_pretrained``. Hugging
+ Face tokenizers are not cached here: configuring one model's template must not
+ change another model's renderer. Hugging Face still caches downloaded files.
+ ART's inference and trajectory code cache these instances where appropriate.
+ Unpinned Llama text models retain Cookbook's public fallback and use its chat
+ template when the base model has none. Explicit revisions keep native files.
+ """
+ # Tinker suffixes identify a model variant, not a tokenizer revision. Local
+ # paths may themselves contain colons.
+ local = os.path.isdir(base_model)
+ model = base_model if local else base_model.split(":", 1)[0]
+ registered = sys.modules.get("tinker_cookbook.tokenizer_utils")
+ is_registered = getattr(registered, "is_tokenizer_registered", None)
+ if callable(is_registered) and is_registered(base_model):
+ assert registered is not None
+ if revision is not None or kwargs:
+ raise ValueError(
+ "Registered tokenizer factories do not accept loader options"
+ )
+ return cast(
+ "PreTrainedTokenizerBase",
+ configure_preserved_thinking_chat_template(
+ copy(registered.get_tokenizer(base_model))
+ ),
+ )
+ if model.startswith("thinkingmachines/Inkling"):
+ if revision is not None or kwargs:
+ raise ValueError("Inkling tokenizers do not accept Hugging Face options")
+ from tinker_cookbook.tokenizer_utils import get_tokenizer as get_tml_tokenizer
+
+ return cast(
+ "PreTrainedTokenizerBase",
+ configure_preserved_thinking_chat_template(
+ copy(get_tml_tokenizer(base_model))
+ ),
+ )
+
+ from transformers import AutoTokenizer, PreTrainedTokenizerFast
+
+ loader: Any = AutoTokenizer
+ if model.startswith("deepseek-ai/DeepSeek-V4-"):
+ loader = PreTrainedTokenizerFast
+ elif model.startswith("moonshotai/Kimi-K2"):
+ # AutoTokenizer can select an incompatible fast backend for Kimi. Its
+ # custom tokenizer also renders tool declarations as TypeScript.
+ if kwargs.get("trust_remote_code") is False:
+ raise ValueError(
+ "Kimi tokenization requires its repository's custom tokenizer code"
+ )
+ from transformers.dynamic_module_utils import get_class_from_dynamic_module
+
+ revision = revision or _KIMI_TOKENIZER_REVISIONS.get(model)
+ loader = cast(
+ "type[PreTrainedTokenizerBase]",
+ get_class_from_dynamic_module(
+ "tokenization_kimi.TikTokenTokenizer",
+ model,
+ revision=revision,
+ **kwargs,
+ ),
+ )
+ kwargs.setdefault(
+ "trust_remote_code",
+ os.path.isdir(model)
+ or os.environ.get("HF_TRUST_REMOTE_CODE", "").lower() in ("1", "true", "yes"),
+ )
+ if revision is not None:
+ kwargs["revision"] = revision
+ try:
+ tokenizer = loader.from_pretrained(model, **kwargs)
+ except OSError:
+ if revision is not None or model not in _LLAMA_TEXT_MODELS:
+ raise
+ # Tinker users need not have access to Meta's gated repositories. Keep
+ # Cookbook's public text-tokenizer fallback, but prefer the actual model
+ # and never apply a model-specific commit to a different repository.
+ tokenizer = loader.from_pretrained(_LLAMA_CHAT_TOKENIZER, **kwargs)
+ if (
+ revision is None
+ and model in _LLAMA_TEXT_MODELS
+ and not getattr(tokenizer, "chat_template", None)
+ ):
+ tokenizer.chat_template = loader.from_pretrained(
+ _LLAMA_CHAT_TOKENIZER, **kwargs
+ ).chat_template
+ if model.startswith("deepseek-ai/DeepSeek-V4-"):
+ from .megatron.dsv4.tokenizer import get_dsv4_tokenizer
+
+ tokenizer = get_dsv4_tokenizer(tokenizer)
+ return cast(
+ "PreTrainedTokenizerBase",
+ configure_preserved_thinking_chat_template(tokenizer),
+ )
diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py
index 007380922..ccd498775 100644
--- a/src/art/trajectories/_tokenize.py
+++ b/src/art/trajectories/_tokenize.py
@@ -24,7 +24,6 @@
from ..utils.chat_template import (
chat_template_with_preserved_thinking,
- configure_preserved_thinking_chat_template,
default_chat_template_kwargs_for_template,
normalize_tool_call_arguments_for_chat_template,
)
@@ -57,9 +56,6 @@
from ._history import _model_matches
from ._protocols import Exchange
-if TYPE_CHECKING:
- from transformers import PreTrainedTokenizerBase
-
_TOKEN_ID = re.compile(r"token_id:(\d+)$")
_WARNED_PREFIX_RETOKENIZATION = False
_TOKENIZER_LOAD_LOCK = threading.Lock()
@@ -1607,22 +1603,10 @@ def _tokenizer_config(model: str, base_model: str | None) -> _TokenizerConfig:
@lru_cache(maxsize=8)
def _cached_tokenizer(base_model: str, revision: str | None) -> Tokenizer:
- try:
- from transformers import AutoTokenizer
- except ImportError as exc:
- raise RuntimeError(
- "Tokenizer fallback requires ART's backend or tinker dependencies"
- ) from exc
- try:
- tokenizer = AutoTokenizer.from_pretrained(
- base_model,
- revision=revision,
- )
- if base_model.startswith("deepseek-ai/DeepSeek-V4-"):
- from ..megatron.dsv4.tokenizer import get_dsv4_tokenizer
+ from ..tokenizer import get_tokenizer
- tokenizer = get_dsv4_tokenizer(cast("PreTrainedTokenizerBase", tokenizer))
- return _as_tokenizer(configure_preserved_thinking_chat_template(tokenizer))
+ try:
+ return _as_tokenizer(get_tokenizer(base_model, revision=revision))
except Exception as exc:
raise ValueError(
f"Could not load tokenizer for {base_model!r}; pass base_model explicitly"
diff --git a/src/art/utils/append_only.py b/src/art/utils/append_only.py
new file mode 100644
index 000000000..114528953
--- /dev/null
+++ b/src/art/utils/append_only.py
@@ -0,0 +1,8 @@
+"""Compatibility imports for ART's lightweight inference helpers."""
+
+import sys
+
+from art_inference import append_only as _implementation
+from art_inference.append_only import * # noqa: F403
+
+sys.modules[__name__] = _implementation
diff --git a/src/art/utils/chat_template.py b/src/art/utils/chat_template.py
index 1be3db827..23e152361 100644
--- a/src/art/utils/chat_template.py
+++ b/src/art/utils/chat_template.py
@@ -1,160 +1,8 @@
-import json
-import re
-from typing import Any
+"""Compatibility imports for ART's lightweight inference helpers."""
-THINKING_CHAT_TEMPLATE_KWARGS: dict[str, Any] = {
- "enable_thinking": False,
- "preserve_thinking": True,
-}
-TOOL_CALL_ARGUMENTS_AS_MAPPING_ATTR = "_art_tool_call_arguments_as_mapping"
-_QWEN_DROP_PRIOR_THINKING = "{%- if loop.index0 > ns.last_query_index %}"
-_QWEN_PRESERVE_PRIOR_THINKING = (
- "{%- if (preserve_thinking is defined and preserve_thinking is true) or "
- "(loop.index0 > ns.last_query_index) %}"
-)
-_GEMMA_DROP_PRIOR_THINKING = (
- "thinking_text and loop.index0 > ns_turn.last_user_idx and "
- "message.get('tool_calls')"
-)
-_GEMMA_PRESERVE_PRIOR_THINKING = (
- "thinking_text and ((preserve_thinking is defined and preserve_thinking is true) "
- "or loop.index0 > ns_turn.last_user_idx) and message.get('tool_calls')"
-)
-_MINIMAX_DROP_PRIOR_THINKING = "reasoning_content and loop.index0 > ns.last_user_index"
-_MINIMAX_PRESERVE_PRIOR_THINKING = (
- "reasoning_content and ((preserve_thinking is defined and preserve_thinking is "
- "true) or loop.index0 > ns.last_user_index)"
-)
+import sys
+from art_inference import chat_template as _implementation
+from art_inference.chat_template import * # noqa: F403
-def chat_template_with_preserved_thinking(chat_template: object) -> object:
- """Add opt-in prior-turn reasoning gates to supported templates."""
- if not isinstance(chat_template, str):
- return chat_template
- replacements = (
- (
- _QWEN_DROP_PRIOR_THINKING,
- _QWEN_PRESERVE_PRIOR_THINKING,
- "enable_thinking" in chat_template,
- ),
- (
- _GEMMA_DROP_PRIOR_THINKING,
- _GEMMA_PRESERVE_PRIOR_THINKING,
- True,
- ),
- (
- _MINIMAX_DROP_PRIOR_THINKING,
- _MINIMAX_PRESERVE_PRIOR_THINKING,
- True,
- ),
- )
- for old, new, supported in replacements:
- if supported and chat_template.count(old) == 1:
- chat_template = chat_template.replace(old, new)
- return chat_template
-
-
-def configure_preserved_thinking_chat_template(tokenizer: object) -> object:
- chat_template = getattr(tokenizer, "chat_template", None)
- configured = chat_template_with_preserved_thinking(chat_template)
- if configured != chat_template:
- setattr(tokenizer, "chat_template", configured)
- return tokenizer
-
-
-def default_chat_template_kwargs_for_template(
- chat_template: object,
-) -> dict[str, Any]:
- kwargs: dict[str, Any] = {}
- if not isinstance(chat_template, str):
- return kwargs
- if "enable_thinking" in chat_template:
- kwargs["enable_thinking"] = False
- if "preserve_thinking" in chat_template:
- kwargs["preserve_thinking"] = True
- if "clear_thinking" in chat_template:
- kwargs["clear_thinking"] = False
- if "deepseek_v4_python_encoder" in chat_template:
- kwargs["drop_thinking"] = False
- return kwargs
-
-
-def default_chat_template_kwargs_for_tokenizer(tokenizer: object) -> dict[str, Any]:
- return default_chat_template_kwargs_for_template(
- getattr(tokenizer, "chat_template", None)
- )
-
-
-def merge_chat_template_kwargs(
- defaults: dict[str, Any] | None,
- overrides: dict[str, Any] | None,
-) -> dict[str, Any]:
- return {**(defaults or {}), **(overrides or {})}
-
-
-def _template_requires_structured_tool_arguments(chat_template: object) -> bool:
- if not isinstance(chat_template, str):
- return False
- arguments_access = (
- r"(?:(? list[dict[str, Any]]:
- """Give chat templates the structured tool arguments they require.
-
- Templates that interpolate the raw JSON string must keep string arguments,
- so only templates that iterate structured arguments trigger normalization.
- """
- if not require_mapping and not _template_requires_structured_tool_arguments(
- chat_template
- ):
- return messages
- normalized: list[dict[str, Any]] = []
- for message in messages:
- calls = message.get("tool_calls")
- if not isinstance(calls, list):
- normalized.append(message)
- continue
- normalized_calls = []
- for call in calls:
- function = call.get("function") if isinstance(call, dict) else None
- arguments = (
- function.get("arguments") if isinstance(function, dict) else None
- )
- if isinstance(arguments, str):
- assert isinstance(function, dict)
- try:
- arguments = json.loads(arguments) if arguments.strip() else {}
- except json.JSONDecodeError as error:
- raise ValueError(
- "tool-call arguments are not valid JSON"
- ) from error
- if not isinstance(arguments, dict):
- raise ValueError("tool-call arguments must decode to a JSON object")
- call = {**call, "function": {**function, "arguments": arguments}}
- normalized_calls.append(call)
- normalized.append({**message, "tool_calls": normalized_calls})
- return normalized
+sys.modules[__name__] = _implementation
diff --git a/src/art_inference/__init__.py b/src/art_inference/__init__.py
new file mode 100644
index 000000000..94378d1a9
--- /dev/null
+++ b/src/art_inference/__init__.py
@@ -0,0 +1 @@
+"""Dependency-free inference helpers shipped with ART and its serving runtime."""
diff --git a/src/art_inference/append_only.py b/src/art_inference/append_only.py
new file mode 100644
index 000000000..989c1c61a
--- /dev/null
+++ b/src/art_inference/append_only.py
@@ -0,0 +1,464 @@
+"""Relate an engine's rendered assistant message to its original sampled IDs.
+
+Only inference may establish these mappings, using its own parsed response and
+renderer. Training tokenization must continue to use the actual served tokens.
+"""
+
+from __future__ import annotations
+
+from collections.abc import Awaitable, Callable, Mapping, Sequence
+from functools import wraps
+import json
+import sys
+from typing import Any
+
+from .token_prefix import PrefixEdit, prefix_edits
+
+PrefixObservation = tuple[list[int], list[int], tuple[PrefixEdit, ...]]
+
+
+def _rendering_edits(tokenizer, rendered, raw):
+ # Stable protocol/multimodal markers separate independent text edits. Keep
+ # them outside replacements so native image offsets remain translatable.
+ special = set(getattr(tokenizer, "all_special_ids", ()))
+ source = [(i, token) for i, token in enumerate(rendered) if token in special]
+ target = [(i, token) for i, token in enumerate(raw) if token in special]
+ if not source or [t for _, t in source] != [t for _, t in target]:
+ return prefix_edits(rendered, raw)
+ edits = []
+ start = raw_start = 0
+ for (stop, _), (raw_stop, _) in zip(
+ [*source, (len(rendered), None)], [*target, (len(raw), None)], strict=True
+ ):
+ edits.extend(
+ PrefixEdit(start + edit.start, start + edit.stop, edit.replacement)
+ for edit in prefix_edits(rendered[start:stop], raw[raw_start:raw_stop])
+ )
+ start, raw_start = stop + 1, raw_stop + 1
+ return tuple(edits)
+
+
+def patch_deepseek_renderer(
+ module: Any, preserves: Callable[[], bool], *, prefix_only: bool = False
+) -> None:
+ """Render historical DeepSeek turns in their original reasoning mode."""
+ render = module.render_message
+
+ @wraps(render)
+ def render_message(index, messages, thinking_mode, *args, **kwargs):
+ if preserves():
+ # A generation setting controls the next answer. Parsed historical
+ # answers already specify whether that turn included reasoning.
+ for message in messages[index:]:
+ if message.get("role") == "assistant":
+ thinking_mode = (
+ "thinking"
+ if any(
+ message.get(name) is not None
+ for name in ("reasoning", "reasoning_content", "thinking")
+ )
+ else "chat"
+ )
+ break
+ if message is not messages[index] and message.get("role") in (
+ "user",
+ "developer",
+ ):
+ break
+ # V3.2 also gates reasoning and the preceding on the last
+ # user in the full chat, even with drop_thinking=False.
+ if prefix_only:
+ messages = messages[: index + 1]
+ message = messages[index]
+ if (
+ message.get("role") == "assistant"
+ and thinking_mode == "thinking"
+ and not any(
+ message.get(name)
+ for name in (
+ "reasoning",
+ "reasoning_content",
+ "thinking",
+ "tool_calls",
+ )
+ )
+ ):
+ # V3.2 rejects an empty reasoning field, although sampling
+ # immediately is valid. Preserve that exact block.
+ return module.thinking_end_token + render(
+ index, messages, "chat", *args, **kwargs
+ )
+ elif not prefix_only:
+ # V4's encoder overrides drop_thinking when tools are present. An
+ # explicit request to remove history still controls each rendering.
+ kwargs["drop_thinking"] = True
+ return render(index, messages, thinking_mode, *args, **kwargs)
+
+ module.render_message = render_message
+
+
+def patch_harmony(module: Any, preserves: Callable[[], bool], namespace: str) -> None:
+ """Keep Harmony analysis and final-turn stops, including imported aliases."""
+ original = module.render_for_completion
+
+ @wraps(original)
+ def render(messages):
+ if not preserves():
+ return original(messages)
+ # Harmony is provided by the inference engine's environment.
+ from openai_harmony import ( # ty: ignore[unresolved-import]
+ Conversation,
+ RenderConversationConfig,
+ Role,
+ )
+
+ encoding = module.get_encoding()
+ config = RenderConversationConfig(auto_drop_analysis=False)
+ # Harmony's whole-conversation renderer changes prior <|return|> to
+ # <|end|>. Rendering each completed message keeps its original framing.
+ tokens = [
+ token
+ for message in messages
+ for token in encoding.render_conversation_for_training(
+ Conversation.from_messages([message]), config=config
+ )
+ ]
+ return tokens + encoding.render_conversation_for_completion(
+ Conversation.from_messages([]), Role.ASSISTANT, config=config
+ )
+
+ replacements = {original: render}
+ drop = getattr(module, "auto_drop_analysis_messages", None)
+ if drop is not None:
+ replacements[drop] = lambda messages: (
+ messages if preserves() else drop(messages)
+ )
+ for name, loaded in tuple(sys.modules.items()):
+ if loaded is not None and name.startswith(namespace):
+ for attribute, value in tuple(vars(loaded).items()):
+ for old, new in replacements.items():
+ if value is old:
+ setattr(loaded, attribute, new)
+ module.render_for_completion = render
+ if drop is not None:
+ module.auto_drop_analysis_messages = replacements[drop]
+
+
+def shifted_span(start: int, length: int, edits: Sequence[PrefixEdit]) -> int:
+ """Move a multimodal span after text edits without changing its contents."""
+ shift = 0
+ for edit in edits:
+ if edit.stop <= start:
+ shift += len(edit.replacement) - (edit.stop - edit.start)
+ elif edit.start < start + length:
+ raise ValueError("History token edit intersects a multimodal placeholder")
+ return start + shift
+
+
+def aligned_values(
+ values: Any, edits: Sequence[PrefixEdit], *, positions=False, fill=None
+):
+ """Translate per-token metadata through edits confined to text tokens."""
+ if positions:
+ import torch
+
+ parts, start, shift = [], 0, 0
+ for edit in edits:
+ parts.append(values[..., start : edit.start] + shift)
+ origin = (
+ values[..., edit.start : edit.start + 1]
+ if edit.start < values.shape[-1]
+ else values[..., -1:] + 1
+ )
+ parts.append(
+ origin + shift + values.new_tensor(list(range(len(edit.replacement))))
+ )
+ shift += len(edit.replacement) - (edit.stop - edit.start)
+ start = edit.stop
+ return torch.cat([*parts, values[..., start:] + shift], dim=-1)
+ result = list(values)
+ for edit in reversed(edits):
+ replaced = result[edit.start : edit.stop]
+ value = (
+ fill
+ if fill is not None
+ else (result[min(edit.start, len(result) - 1)] if result else 0)
+ )
+ if any(item != value for item in replaced):
+ raise ValueError("History token edit crosses a per-token metadata boundary")
+ result[edit.start : edit.stop] = [value] * len(edit.replacement)
+ return values.new_tensor(result) if hasattr(values, "new_tensor") else result
+
+
+def output_prefix_observations(
+ tokenizer: Any,
+ rendered_prompt: Sequence[int],
+ raw_prompt: Sequence[int],
+ raw_output: Sequence[int],
+ *,
+ prompt_edits: Sequence[PrefixEdit] | None = None,
+) -> list[PrefixObservation]:
+ """Preserve native IDs for the response and its last reasoning boundary.
+
+ In particular, reasoning can survive a later action edit, for streaming and
+ non-streaming generations in any protocol. A decode/encode pass must retain
+ the same special-token sequence before we use an intermediate boundary.
+ Keep at most two entries so tool-heavy outputs do not multiply prompt storage.
+ """
+ output = list(raw_output)
+ if not output:
+ return []
+ rendered = list(
+ tokenizer.encode(_decode(tokenizer, output), add_special_tokens=False)
+ )
+ if not rendered:
+ return []
+ special_ids = set(getattr(tokenizer, "all_special_ids", ()))
+ raw_boundaries = [
+ (i + 1, token) for i, token in enumerate(output) if token in special_ids
+ ]
+ rendered_boundaries = [
+ (i + 1, token) for i, token in enumerate(rendered) if token in special_ids
+ ]
+ boundaries = []
+ if [token for _, token in raw_boundaries] == [
+ token for _, token in rendered_boundaries
+ ]:
+ boundaries = [
+ (r[0], n[0])
+ for r, n in zip(rendered_boundaries, raw_boundaries, strict=True)
+ if _decode(tokenizer, [r[1]]) in {"", "", "<|end|>"}
+ ][-1:]
+ if not boundaries or boundaries[-1] != (len(rendered), len(output)):
+ boundaries.append((len(rendered), len(output)))
+ entries = []
+ prompt_edits = (
+ tuple(prompt_edits)
+ if prompt_edits is not None
+ else _rendering_edits(tokenizer, rendered_prompt, raw_prompt)
+ )
+ offset = len(rendered_prompt)
+ for rendered_end, raw_end in boundaries:
+ edits = tuple(
+ PrefixEdit(offset + edit.start, offset + edit.stop, edit.replacement)
+ for edit in _rendering_edits(
+ tokenizer, rendered[:rendered_end], output[:raw_end]
+ )
+ )
+ entries.append(
+ (
+ [*rendered_prompt, *rendered[:rendered_end]],
+ [*raw_prompt, *output[:raw_end]],
+ (*prompt_edits, *edits),
+ )
+ )
+ return entries
+
+
+def merge_chat_delta(message: dict[str, Any], delta: dict[str, Any]) -> None:
+ for name, value in delta.items():
+ if name == "tool_calls" and isinstance(value, list):
+ calls = message.setdefault("tool_calls", [])
+ for part in value:
+ index = part["index"]
+ while len(calls) <= index:
+ calls.append({})
+ merge_chat_delta(
+ calls[index], {k: v for k, v in part.items() if k != "index"}
+ )
+ elif isinstance(value, dict):
+ merge_chat_delta(message.setdefault(name, {}), value)
+ elif isinstance(value, str):
+ message[name] = (
+ value
+ if name in {"role", "type", "id"}
+ else message.get(name, "") + value
+ )
+
+
+def preserves_history(kwargs: Mapping[str, Any] | None) -> bool:
+ options = kwargs or {}
+ return not (
+ options.get("preserve_thinking") is False
+ or options.get("clear_thinking") is True
+ or options.get("drop_thinking") is True
+ )
+
+
+def has_renderable_tool_arguments(message: Mapping[str, Any]) -> bool:
+ functions = [call.get("function", {}) for call in message.get("tool_calls") or []]
+ if message.get("function_call"):
+ functions.append(message["function_call"])
+ try:
+ return all(
+ not isinstance(f.get("arguments"), str)
+ or isinstance(json.loads(f["arguments"] or "{}"), dict)
+ for f in functions
+ )
+ except json.JSONDecodeError:
+ return False
+
+
+def _decode(tokenizer: Any, tokens: Sequence[int]) -> str:
+ return tokenizer.decode(list(tokens), skip_special_tokens=False)
+
+
+def _whitespace_prefix(tokenizer: Any, rendered: list[int], raw: list[int]) -> int:
+ """Find an exact raw-token boundary for a whitespace-normalized prefix."""
+ target = _decode(tokenizer, rendered)
+ source = _decode(tokenizer, raw)
+ i = j = 0
+ while i < len(target):
+ if target[i].isspace():
+ i += 1
+ elif j < len(source) and source[j].isspace():
+ j += 1
+ elif j < len(source) and target[i] == source[j]:
+ i += 1
+ j += 1
+ else:
+ return 0
+ if target and target[-1].isspace():
+ while j < len(source) and source[j].isspace():
+ j += 1
+ # Decoding keeps the native IDs, including non-canonical tokenizations. A
+ # boundary inside a token cannot be preserved independently of the action.
+ low, high = 0, len(raw)
+ while low < high:
+ middle = (low + high) // 2
+ if len(_decode(tokenizer, raw[:middle])) < j:
+ low = middle + 1
+ else:
+ high = middle
+ return low if _decode(tokenizer, raw[:low]) == source[:j] else 0
+
+
+def chat_prefix_observations(
+ tokenizer: Any,
+ rendered_prompt: Sequence[int],
+ completed_prompt: Sequence[int],
+ raw_prompt: Sequence[int],
+ raw_output: Sequence[int],
+ *,
+ reasoning_prompt: Sequence[int] | None = None,
+ complete: bool = True,
+) -> list[PrefixObservation]:
+ """Record a native response and, separately, its unchanged reasoning.
+
+ ``completed_prompt`` must come from rendering the original request plus the
+ engine's parsed response. ``reasoning_prompt`` renders the same request with
+ only that response's reasoning. Keeping the latter mapping lets callers edit
+ the action without losing the exact reasoning prefix. Explicit requests to
+ drop thinking must bypass observation and lookup.
+
+ Truncated responses may contribute a reasoning prefix, but cannot replace a
+ completed turn's framing. No sampled IDs or log probabilities are invented.
+ """
+ prompt, completed = list(rendered_prompt), list(completed_prompt)
+ output = list(raw_output)
+ if not output or completed[: len(prompt)] != prompt:
+ return []
+ entries: list[PrefixObservation] = []
+ if reasoning_prompt is not None:
+ boundary = next(
+ (i for i, (a, b) in enumerate(zip(completed, reasoning_prompt)) if a != b),
+ min(len(completed), len(reasoning_prompt)),
+ )
+ if boundary > len(prompt):
+ sampled_boundary = _whitespace_prefix(
+ tokenizer, completed[len(prompt) : boundary], output
+ )
+ if sampled_boundary:
+ rendered = completed[:boundary]
+ raw = [*raw_prompt, *output[:sampled_boundary]]
+ entries.append(
+ (rendered, raw, _rendering_edits(tokenizer, rendered, raw))
+ )
+ if complete and len(completed) > len(prompt):
+ # A template's separator after the sampled stop belongs to the next
+ # turn. Leave it in the rendered suffix instead of swallowing it.
+ while (
+ len(completed) > len(prompt)
+ and _decode(tokenizer, completed[-1:]).isspace()
+ ):
+ completed.pop()
+ # User stop strings (and APIs omitting the sampled stop ID) do not
+ # prove the template's terminal framing. Never delete that framing.
+ if completed[-1] != output[-1]:
+ return entries
+ raw = [*raw_prompt, *output]
+ entries.append((completed, raw, _rendering_edits(tokenizer, completed, raw)))
+ return entries
+
+
+async def chat_response_prefixes(
+ tokenizer: Any,
+ request: Any,
+ raw_prompt: Sequence[int],
+ choices: Sequence[tuple[Mapping[str, Any], Sequence[int], bool]],
+ render: Callable[[Any], Awaitable[list[int] | None]],
+) -> list[PrefixObservation]:
+ """Observe parsed native chat choices using that engine's request renderer."""
+ payload = request.model_dump(mode="python")
+ if not preserves_history(payload.get("chat_template_kwargs")):
+ return []
+ rendered = await render(request)
+ messages = payload.get("messages")
+ if rendered is None or not isinstance(messages, list):
+ return []
+ fields = type(request).model_fields
+ payload.update(
+ (name, value)
+ for name, value in (
+ ("add_generation_prompt", False),
+ ("continue_final_message", False),
+ ("input_ids", None),
+ ("n", 1),
+ ("stream", False),
+ ("stream_options", None),
+ # These turns have already completed: reserve no further output and
+ # retain the entire history when observing its rendered tokens.
+ ("max_tokens", None),
+ ("max_completion_tokens", None),
+ ("truncate_prompt_tokens", None),
+ )
+ if name in fields
+ )
+
+ async def complete(message: Mapping[str, Any]) -> list[int] | None:
+ return await render(
+ type(request).model_validate({**payload, "messages": [*messages, message]})
+ )
+
+ entries: list[PrefixObservation] = []
+ for message, output, finished in choices:
+ # An invalid sampled call must reach the caller unchanged. It cannot be
+ # rendered as a completed tool turn; native special-token observations
+ # still preserve its reasoning prefix.
+ if not has_renderable_tool_arguments(message):
+ continue
+ completed = await complete(message)
+ if completed is None:
+ continue
+ reasoning = {
+ key: message[key]
+ for key in ("reasoning", "reasoning_content", "thinking")
+ if message.get(key) is not None
+ }
+ reasoning_prompt = (
+ await complete({"role": "assistant", "content": "", **reasoning})
+ if reasoning
+ else None
+ )
+ entries.extend(
+ chat_prefix_observations(
+ tokenizer,
+ rendered,
+ completed,
+ raw_prompt,
+ output,
+ reasoning_prompt=reasoning_prompt,
+ complete=finished,
+ )
+ )
+ return entries
diff --git a/src/art_inference/chat_template.py b/src/art_inference/chat_template.py
new file mode 100644
index 000000000..20dd21c7e
--- /dev/null
+++ b/src/art_inference/chat_template.py
@@ -0,0 +1,258 @@
+import json
+import re
+from typing import Any
+
+THINKING_CHAT_TEMPLATE_KWARGS: dict[str, Any] = {
+ "enable_thinking": False,
+ "preserve_thinking": True,
+}
+TOOL_CALL_ARGUMENTS_AS_MAPPING_ATTR = "_art_tool_call_arguments_as_mapping"
+_QWEN_DROP_PRIOR_THINKING = "{%- if loop.index0 > ns.last_query_index %}"
+_QWEN_PRESERVE_PRIOR_THINKING = (
+ "{%- if (preserve_thinking is defined and preserve_thinking is true) or "
+ "(loop.index0 > ns.last_query_index) %}"
+)
+_GEMMA_DROP_PRIOR_THINKING = (
+ "thinking_text and loop.index0 > ns_turn.last_user_idx and "
+ "message.get('tool_calls')"
+)
+_GEMMA_PRESERVE_PRIOR_THINKING = (
+ "thinking_text and ((preserve_thinking is defined and preserve_thinking is true) "
+ "or (loop.index0 > ns_turn.last_user_idx and message.get('tool_calls')))"
+)
+_MINIMAX_DROP_PRIOR_THINKING = "reasoning_content and loop.index0 > ns.last_user_index"
+_MINIMAX_PRESERVE_PRIOR_THINKING = (
+ "reasoning_content and ((preserve_thinking is defined and preserve_thinking is "
+ "true) or loop.index0 > ns.last_user_index)"
+)
+
+
+def chat_template_with_preserved_thinking(chat_template: object) -> object:
+ """Preserve prior reasoning by default, while respecting explicit opt-outs."""
+ if isinstance(chat_template, dict):
+ return {
+ name: chat_template_with_preserved_thinking(template)
+ for name, template in chat_template.items()
+ }
+ if not isinstance(chat_template, str):
+ return chat_template
+ replacements = (
+ (
+ _QWEN_DROP_PRIOR_THINKING,
+ _QWEN_PRESERVE_PRIOR_THINKING,
+ "enable_thinking" in chat_template,
+ ),
+ (
+ _GEMMA_DROP_PRIOR_THINKING,
+ _GEMMA_PRESERVE_PRIOR_THINKING,
+ True,
+ ),
+ (
+ _MINIMAX_DROP_PRIOR_THINKING,
+ _MINIMAX_PRESERVE_PRIOR_THINKING,
+ True,
+ ),
+ (
+ "(loop.index0 > ns_turn.last_user_idx) or (preserve_thinking and message.get('tool_calls'))",
+ "preserve_thinking or (loop.index0 > ns_turn.last_user_idx)",
+ True,
+ ),
+ *(
+ (
+ f"message.{field} and not future_final_message.found",
+ f"message.{field} and ((preserve_thinking is defined and preserve_thinking is true) or not future_final_message.found)",
+ True,
+ )
+ for field in ("content", "thinking")
+ ),
+ (
+ "{#- CoT is dropped during all previous turns, so we never render it for inference #}",
+ "{%- if preserve_thinking is defined and preserve_thinking is true and message.thinking is defined %}<|start|>assistant<|channel|>analysis<|message|>{{ message.thinking }}<|end|>{%- endif %}",
+ True,
+ ),
+ (
+ '"<|start|>assistant<|channel|>final<|message|>" + message.content + "<|end|>"',
+ '"<|start|>assistant<|channel|>final<|message|>" + message.content + ("<|return|>" if preserve_thinking else "<|end|>")',
+ True,
+ ),
+ )
+ for old, new, supported in replacements:
+ if supported and chat_template.count(old) == 1:
+ chat_template = chat_template.replace(old, new)
+ # Kimi 2.5 splits historical turns into a branch that blanks reasoning.
+ # Preserve them using its ordinary assistant-turn branch instead.
+ if (
+ "set hist_msgs = messages[:ns.last_non_tool_call_assistant_msg+1]"
+ in chat_template
+ ):
+ chat_template = (
+ chat_template.replace(
+ "set hist_msgs = messages[:ns.last_non_tool_call_assistant_msg+1]",
+ "set hist_msgs = [] if preserve_thinking else messages[:ns.last_non_tool_call_assistant_msg+1]",
+ )
+ .replace(
+ "set suffix_msgs = messages[ns.last_non_tool_call_assistant_msg+1:]",
+ "set suffix_msgs = messages if preserve_thinking else messages[ns.last_non_tool_call_assistant_msg+1:]",
+ )
+ .replace(
+ "{%- if thinking is defined and thinking is false -%}",
+ "{%- if thinking is defined and thinking is false and not preserve_thinking -%}",
+ 1,
+ )
+ )
+ # Qwen's native reasoning parser returns the sampled whitespace. The stock
+ # template trims it, then invents a newline before . Preserve the
+ # parsed field verbatim; keep the original non-thinking/legacy fallback.
+ if (
+ "" in chat_template
+ and "reasoning_content|trim" in chat_template
+ and "if not preserve_thinking or message.reasoning_content" not in chat_template
+ ):
+ chat_template = chat_template.replace(
+ "{%- set reasoning_content = reasoning_content|trim %}",
+ "{%- if not preserve_thinking or message.reasoning_content is not string %}"
+ "{%- set reasoning_content = reasoning_content|trim %}{%- endif %}",
+ )
+ chat_template = chat_template.replace(
+ "reasoning_content + '\\n\\n\\n'",
+ "reasoning_content + ('\\n\\n' if preserve_thinking and message.reasoning_content is string and reasoning_content else '\\n\\n\\n')",
+ )
+ chat_template = chat_template.replace(
+ "set content = render_content(message.content, true)|trim",
+ "set content = (render_content(message.content, true) if preserve_thinking and message.role == 'assistant' else render_content(message.content, true)|trim)",
+ )
+ if "clear_thinking" in chat_template:
+ chat_template = chat_template.replace(
+ "{{ content.strip() }}",
+ "{{ content if not clear_thinking else content.strip() }}",
+ ).replace(
+ "{%- if content.strip() -%}",
+ "{%- if (content if not clear_thinking else content.strip()) -%}",
+ )
+ if "ns_turn.last_user_idx" in chat_template:
+ for value in ("message['content']", "item['text']"):
+ chat_template = chat_template.replace(
+ f"{{{{- {value} | trim -}}}}",
+ f"{{{{- {value} if preserve_thinking and message.role == 'assistant' else {value} | trim -}}}}",
+ )
+ # Embed preservation defaults so direct tokenizer use and inference engines
+ # agree with ART. Passing a value explicitly still takes precedence.
+ for name, value in (("preserve_thinking", "true"), ("clear_thinking", "false")):
+ default = f"{{%- set {name} = {name} | default({value}) -%}}"
+ if name in chat_template and default not in chat_template:
+ chat_template = default + chat_template
+ return chat_template
+
+
+def configure_preserved_thinking_chat_template(tokenizer: object) -> object:
+ chat_template = getattr(tokenizer, "chat_template", None)
+ configured = chat_template_with_preserved_thinking(chat_template)
+ if configured != chat_template:
+ setattr(tokenizer, "chat_template", configured)
+ return tokenizer
+
+
+def default_chat_template_kwargs_for_template(
+ chat_template: object,
+) -> dict[str, Any]:
+ kwargs = default_preservation_kwargs_for_template(chat_template)
+ if not isinstance(chat_template, str):
+ return kwargs
+ if "enable_thinking" in chat_template:
+ kwargs["enable_thinking"] = False
+ return kwargs
+
+
+def default_preservation_kwargs_for_template(chat_template: object) -> dict[str, Any]:
+ """History defaults independent of whether the next turn should think."""
+ kwargs: dict[str, Any] = {}
+ if not isinstance(chat_template, str):
+ return kwargs
+ if "preserve_thinking" in chat_template:
+ kwargs["preserve_thinking"] = True
+ if "clear_thinking" in chat_template:
+ kwargs["clear_thinking"] = False
+ if "deepseek_v4_python_encoder" in chat_template:
+ kwargs["drop_thinking"] = False
+ return kwargs
+
+
+def default_chat_template_kwargs_for_tokenizer(tokenizer: object) -> dict[str, Any]:
+ return default_chat_template_kwargs_for_template(
+ getattr(tokenizer, "chat_template", None)
+ )
+
+
+def merge_chat_template_kwargs(
+ defaults: dict[str, Any] | None,
+ overrides: dict[str, Any] | None,
+) -> dict[str, Any]:
+ return {**(defaults or {}), **(overrides or {})}
+
+
+def _template_requires_structured_tool_arguments(chat_template: object) -> bool:
+ if not isinstance(chat_template, str):
+ return False
+ arguments_access = (
+ r"(?:(? list[dict[str, Any]]:
+ """Give chat templates the structured tool arguments they require.
+
+ Templates that interpolate the raw JSON string must keep string arguments,
+ so only templates that iterate structured arguments trigger normalization.
+ """
+ if not require_mapping and not _template_requires_structured_tool_arguments(
+ chat_template
+ ):
+ return messages
+ normalized: list[dict[str, Any]] = []
+ for message in messages:
+ calls = message.get("tool_calls")
+ if not isinstance(calls, list):
+ normalized.append(message)
+ continue
+ normalized_calls = []
+ for call in calls:
+ function = call.get("function") if isinstance(call, dict) else None
+ arguments = (
+ function.get("arguments") if isinstance(function, dict) else None
+ )
+ if isinstance(arguments, str):
+ assert isinstance(function, dict)
+ try:
+ arguments = json.loads(arguments) if arguments.strip() else {}
+ except json.JSONDecodeError as error:
+ raise ValueError(
+ "tool-call arguments are not valid JSON"
+ ) from error
+ if not isinstance(arguments, dict):
+ raise ValueError("tool-call arguments must decode to a JSON object")
+ call = {**call, "function": {**function, "arguments": arguments}}
+ normalized_calls.append(call)
+ normalized.append({**message, "tool_calls": normalized_calls})
+ return normalized
diff --git a/src/art_inference/sglang.py b/src/art_inference/sglang.py
new file mode 100644
index 000000000..680a9bd39
--- /dev/null
+++ b/src/art_inference/sglang.py
@@ -0,0 +1,382 @@
+"""History preservation for SGLang's native chat, Responses, and Messages APIs."""
+
+from __future__ import annotations
+
+from collections.abc import AsyncIterator, Awaitable, Callable
+from contextvars import ContextVar
+import copy
+from functools import wraps
+import importlib
+import json
+from typing import Any
+
+from .append_only import (
+ aligned_values,
+ chat_response_prefixes,
+ merge_chat_delta,
+ patch_deepseek_renderer,
+ patch_harmony,
+ preserves_history,
+ shifted_span,
+)
+from .chat_template import (
+ chat_template_with_preserved_thinking,
+ configure_preserved_thinking_chat_template,
+)
+from .token_prefix import apply_prefix_edits
+
+_DROP_THINKING: ContextVar[bool] = ContextVar("art_sglang_drop_thinking", default=False)
+
+
+def replace_prompt_tokens(tokenized, edits):
+ """Translate expanded multimodal positions with the preserved text prefix."""
+ if not edits:
+ return
+ mm = getattr(tokenized, "mm_inputs", None)
+ if mm is not None:
+ mm = copy.copy(mm)
+ mm.mm_items = [copy.copy(item) for item in mm.mm_items]
+ for item in mm.mm_items:
+ if item.offsets is not None:
+ item.offsets = [
+ (
+ shifted_span(start, end - start + 1, edits),
+ shifted_span(start, end - start + 1, edits) + end - start,
+ )
+ for start, end in item.offsets
+ ]
+ for name in ("input_ids", "padded_input_ids"):
+ if getattr(mm, name, None) is not None:
+ setattr(mm, name, apply_prefix_edits(list(getattr(mm, name)), edits))
+ for name in ("mrope_positions", "token_type_ids", "visible_frame_counts"):
+ if getattr(mm, name, None) is not None:
+ setattr(
+ mm,
+ name,
+ aligned_values(
+ getattr(mm, name),
+ edits,
+ positions=name == "mrope_positions",
+ fill=0 if name == "token_type_ids" else None,
+ ),
+ )
+ tokenized.mm_inputs = mm
+ if getattr(tokenized, "token_type_ids", None) is not None:
+ tokenized.token_type_ids = aligned_values(
+ tokenized.token_type_ids, edits, fill=0
+ )
+
+
+class _TokenizerView:
+ def __init__(self, tokenizer: Any, generation: bool):
+ self._tokenizer = tokenizer
+ self._generation = generation
+
+ def __getattr__(self, name):
+ return getattr(self._tokenizer, name)
+
+ def apply_chat_template(self, *args, **kwargs):
+ kwargs["add_generation_prompt"] = self._generation
+ return self._tokenizer.apply_chat_template(*args, **kwargs)
+
+
+def patch_history(
+ observe: Callable[[Any, list[Any]], Awaitable[None]],
+ importer=importlib.import_module,
+) -> None:
+ chat = importer("sglang.srt.entrypoints.openai.serving_chat").OpenAIServingChat
+ responses = importer(
+ "sglang.srt.entrypoints.openai.serving_responses"
+ ).OpenAIServingResponses
+ if getattr(chat, "_art_append_only", False):
+ return
+ patch_harmony(
+ importer("sglang.srt.entrypoints.harmony_utils"),
+ lambda: not _DROP_THINKING.get(),
+ "sglang.",
+ )
+ process = chat._process_messages
+
+ def patch_encoder(encoding):
+ encode = encoding.encode_messages
+
+ @wraps(encode)
+ def encode_messages(*args, **kwargs):
+ kwargs.setdefault("drop_thinking", _DROP_THINKING.get())
+ return encode(*args, **kwargs)
+
+ encoding.encode_messages = encode_messages
+
+ for name in ("dsv4", "dsv32"):
+ encoding = importer(f"sglang.srt.entrypoints.openai.encoding_{name}")
+ patch_encoder(encoding)
+ patch_deepseek_renderer(
+ encoding, lambda: not _DROP_THINKING.get(), prefix_only=name == "dsv32"
+ )
+
+ @wraps(process)
+ def process_messages(self, request, is_multimodal):
+ configure_preserved_thinking_chat_template(self.tokenizer_manager.tokenizer)
+ if getattr(request, "chat_template", None) is not None:
+ request.chat_template = chat_template_with_preserved_thinking(
+ request.chat_template
+ )
+ # Templates already embed their defaults. Adding them to this request
+ # would shadow a Responses adapter's outer rendering options. Likewise,
+ # retain its encoder context unless this chat request overrides it.
+ options = getattr(request, "chat_template_kwargs", None) or {}
+ token = _DROP_THINKING.set(
+ not preserves_history(options)
+ if any(
+ k in options
+ for k in ("preserve_thinking", "clear_thinking", "drop_thinking")
+ )
+ else _DROP_THINKING.get()
+ )
+ try:
+ return process(self, request, is_multimodal)
+ finally:
+ _DROP_THINKING.reset(token)
+
+ chat._process_messages = process_messages
+
+ async def record(self, request, raw_request, prompt, choices):
+ if prompt is None or not preserves_history(
+ getattr(request, "chat_template_kwargs", None)
+ ):
+ return
+
+ async def render(value):
+ view = copy.copy(self)
+ generation = len(value.messages) == len(request.messages)
+ if not generation:
+ # Native request rendering converts a trailing assistant into
+ # a user turn. This view renders an already completed response.
+ view._handle_last_assistant_message = lambda messages, _: (
+ messages,
+ None,
+ )
+ view.tokenizer_manager = copy.copy(self.tokenizer_manager)
+ view.tokenizer_manager.tokenizer = _TokenizerView(
+ self.tokenizer_manager.tokenizer,
+ generation,
+ )
+ result = view._process_messages(
+ value,
+ is_multimodal=self.tokenizer_manager.model_config.is_multimodal,
+ ).prompt_ids
+ return result if isinstance(result, list) and result else None
+
+ entries = await chat_response_prefixes(
+ self.tokenizer_manager.tokenizer, request, prompt, choices, render
+ )
+ if entries:
+ await observe(raw_request, entries)
+
+ full = chat._handle_non_streaming_request
+
+ @wraps(full)
+ async def complete(self, adapted_request, request, raw_request):
+ result = await full(self, adapted_request, request, raw_request)
+ payload = (
+ result.model_dump(exclude_none=True)
+ if hasattr(result, "model_dump")
+ else json.loads(result.body)
+ )
+ prompt = payload.get("prompt_token_ids")
+ choices = []
+ for choice in payload.get("choices", []):
+ prompt = prompt or choice.get("prompt_token_ids")
+ ids = choice.get("token_ids") or (choice.get("logprobs") or {}).get(
+ "token_ids"
+ )
+ if ids is not None:
+ choices.append(
+ (
+ choice["message"],
+ ids,
+ choice.get("finish_reason") not in ("length", "abort"),
+ )
+ )
+ if choices:
+ await record(self, request, raw_request, prompt, choices)
+ return result
+
+ chat._handle_non_streaming_request = complete
+ generate = chat._generate_chat_stream
+
+ @wraps(generate)
+ async def stream(self, adapted_request, request, raw_request):
+ messages, tokens, finished = {}, {}, {}
+ prompt = None
+ async for event in generate(self, adapted_request, request, raw_request):
+ if isinstance(event, str) and event.startswith("data: "):
+ data = event[6:].strip()
+ if data == "[DONE]":
+ await record(
+ self,
+ request,
+ raw_request,
+ prompt,
+ [
+ (message, tokens[index], finished.get(index, False))
+ for index, message in messages.items()
+ if index in tokens
+ ],
+ )
+ else:
+ payload = json.loads(data)
+ prompt = prompt or payload.get("prompt_token_ids")
+ for choice in payload.get("choices", []):
+ index = choice["index"]
+ merge_chat_delta(
+ messages.setdefault(index, {"role": "assistant"}),
+ choice.get("delta") or {},
+ )
+ if "token_ids" in choice:
+ tokens.setdefault(index, []).extend(choice["token_ids"])
+ if choice.get("finish_reason"):
+ finished[index] = choice["finish_reason"] not in (
+ "length",
+ "abort",
+ )
+ yield event
+
+ chat._generate_chat_stream = stream
+
+ normalize = responses._normalize_response_message_for_chat
+
+ @classmethod
+ def normalize_message(cls, message):
+ if hasattr(message, "model_dump"):
+ message = message.model_dump(exclude_none=True)
+ if isinstance(message, dict) and message.get("type") == "reasoning":
+ # A summary may differ from the sampled reasoning. Joining content
+ # parts with newlines would also alter the original response.
+ parts = message.get("content") or message.get("summary") or []
+ text = "".join(part.get("text", "") for part in parts)
+ return {"role": "assistant", "reasoning_content": text} if text else None
+ return normalize(message)
+
+ responses._normalize_response_message_for_chat = normalize_message
+ construct = responses._construct_input_messages
+
+ @wraps(construct)
+ def input_messages(self, request, prev_response=None):
+ if prev_response is not None:
+ inputs = (
+ [{"role": "user", "content": request.input}]
+ if isinstance(request.input, str)
+ else request.input
+ )
+ request = request.model_copy(
+ update={"input": [*prev_response.output, *inputs]}
+ )
+ prev_response = prev_response.model_copy(update={"output": []})
+ return construct(self, request, prev_response)
+
+ responses._construct_input_messages = input_messages
+ harmony = responses._construct_input_messages_with_harmony
+
+ @wraps(harmony)
+ def harmony_messages(self, request, prev_response):
+ if prev_response is None:
+ return harmony(self, request, prev_response)
+ # SGLang deletes reasoning from the stored list in place. Give it a
+ # private list, and retain the original prefix unless rewriting was
+ # explicitly requested for this call.
+ previous = list(self.msg_store[prev_response.id])
+ view = copy.copy(self)
+ view.msg_store = {**self.msg_store, prev_response.id: previous.copy()}
+ result = harmony(view, request, prev_response)
+ retained = view.msg_store[prev_response.id]
+ if preserves_history(getattr(request, "chat_template_kwargs", None)):
+ if result[: len(retained)] != retained:
+ raise RuntimeError("SGLang's stored-history rendering contract changed")
+ return [*previous, *result[len(retained) :]]
+ return result
+
+ responses._construct_input_messages_with_harmony = harmony_messages
+
+ create = responses.create_responses
+
+ @wraps(create)
+ async def create_responses(self, request, raw_request=None):
+ async def record_response(payload):
+ generations = payload.get("token_generations") or []
+ if getattr(self, "use_harmony", False) or len(generations) != 1:
+ return
+ generation = generations[0]
+ messages = self._merge_consecutive_assistant_messages(
+ [
+ message
+ for item in payload.get("output", [])
+ if (message := self._normalize_response_message_for_chat(item))
+ is not None
+ ]
+ )
+ if len(messages) != 1 or messages[0].get("role") != "assistant":
+ return
+ previous = (
+ self.response_store.get(request.previous_response_id)
+ if request.previous_response_id
+ else None
+ )
+ protocol = importer("sglang.srt.entrypoints.openai.protocol")
+ view = protocol.ChatCompletionRequest(
+ model=request.model,
+ messages=self._construct_input_messages(request, previous),
+ tools=self._response_tools_to_chat_tools(request) or None,
+ chat_template=getattr(request, "chat_template", None),
+ chat_template_kwargs=getattr(request, "chat_template_kwargs", None),
+ )
+ await record(
+ self,
+ view,
+ raw_request,
+ generation["prompt_token_ids"],
+ [
+ (
+ messages[0],
+ [token["token_id"] for token in generation["output_tokens"]],
+ payload.get("status") == "completed",
+ )
+ ],
+ )
+
+ token = _DROP_THINKING.set(
+ not preserves_history(getattr(request, "chat_template_kwargs", None))
+ )
+ try:
+ result = await create(self, request, raw_request)
+ finally:
+ _DROP_THINKING.reset(token)
+ if not isinstance(result, AsyncIterator):
+ payload = (
+ result.model_dump(exclude_none=True)
+ if hasattr(result, "model_dump")
+ else json.loads(result.body)
+ )
+ await record_response(payload)
+ return result
+
+ async def stream_response():
+ token = _DROP_THINKING.set(
+ not preserves_history(getattr(request, "chat_template_kwargs", None))
+ )
+ try:
+ async for event in result:
+ if isinstance(event, str) and event.startswith(
+ ("event: response.completed\n", "event: response.incomplete\n")
+ ):
+ await record_response(
+ json.loads(event.split("data: ", 1)[1])["response"]
+ )
+ yield event
+ finally:
+ _DROP_THINKING.reset(token)
+
+ return stream_response()
+
+ responses.create_responses = create_responses
+ chat._art_append_only = True
diff --git a/src/art_inference/token_prefix.py b/src/art_inference/token_prefix.py
new file mode 100644
index 000000000..9fdd8aff4
--- /dev/null
+++ b/src/art_inference/token_prefix.py
@@ -0,0 +1,1054 @@
+"""Bounded mappings from rendered history to previously served token IDs."""
+
+from __future__ import annotations
+
+from array import array
+import asyncio
+import base64
+import binascii
+from collections import OrderedDict
+from collections.abc import Awaitable, Callable
+from contextlib import suppress
+from dataclasses import dataclass, field
+import gzip
+import hashlib
+import json
+import logging
+import string
+import sys
+import time
+from typing import Sequence
+import zlib
+
+logger = logging.getLogger(__name__)
+
+COMPACT_PREFIX_VERSION = "sha256-u32be-v1"
+COMPACT_PREFIX_MAX_CANDIDATES = 1024
+COMPACT_PREFIX_MAX_REPLACEMENT_IDS = 262_144
+COMPACT_PREFIX_RECORD_LIMIT = 4 * 1024 * 1024
+COMPACT_PREFIX_RESPONSE_LIMIT = 8 * 1024 * 1024
+COMPACT_PREFIX_HEADER_LIMIT = 6 * 1024
+COMPACT_PREFIX_HEADER_DECODE_LIMIT = 64 * 1024
+# Keep the original wire domain so deployed prefix stores remain compatible.
+_DIGEST_DOMAIN = b"caladan-token-prefix\0" + COMPACT_PREFIX_VERSION.encode() + b"\0"
+_MAX_PREFIX_TOKENS = 262_144
+_SharedPrefixEntry = tuple[str, list[int], list[int], str, tuple["PrefixEdit", ...]]
+
+
+@dataclass(frozen=True, slots=True)
+class PrefixEdit:
+ start: int
+ stop: int
+ replacement: tuple[int, ...]
+
+
+@dataclass(frozen=True, slots=True)
+class CompactPrefixCandidate:
+ rendered_length: int
+ rendered_digest: str
+ raw_digest: str
+ edits: tuple[PrefixEdit, ...]
+
+
+@dataclass(frozen=True, slots=True)
+class PrefixMatch:
+ rendered_length: int
+ raw_prefix: tuple[int, ...]
+ edits: tuple[PrefixEdit, ...] = field(default=(), compare=False)
+
+
+def _token_bytes(token_ids: Sequence[int]) -> bytes:
+ # Engine and wire boundaries validate bool/type explicitly. Let the C array
+ # conversion enforce the numeric/range constraint without another O(n)
+ # Python scan on every 128K lookup.
+ try:
+ values = array("I", token_ids)
+ except (OverflowError, TypeError, ValueError) as error:
+ raise ValueError("token IDs must be unsigned 32-bit integers") from error
+ if values.itemsize != 4:
+ raise RuntimeError("this platform does not provide 32-bit unsigned integers")
+ if sys.byteorder == "little":
+ values.byteswap()
+ return values.tobytes()
+
+
+def token_prefix_digests(
+ token_ids: Sequence[int], lengths: Sequence[int]
+) -> dict[int, str]:
+ """Hash several token prefixes in one pass over a fixed-width encoding."""
+ requested = sorted(set(lengths))
+ if any(
+ not isinstance(length, int)
+ or isinstance(length, bool)
+ or not 0 <= length <= len(token_ids)
+ for length in requested
+ ):
+ raise ValueError("invalid token prefix length")
+ if not requested:
+ return {}
+ encoded = _token_bytes(token_ids[: requested[-1]])
+ digest = hashlib.sha256(_DIGEST_DOMAIN)
+ previous = 0
+ result: dict[int, str] = {}
+ for length in requested:
+ digest.update(encoded[previous * 4 : length * 4])
+ result[length] = digest.hexdigest()
+ previous = length
+ return result
+
+
+def token_digest(token_ids: Sequence[int]) -> str:
+ return token_prefix_digests(token_ids, [len(token_ids)])[len(token_ids)]
+
+
+def apply_prefix_edits(
+ rendered_prefix: Sequence[int], edits: Sequence[PrefixEdit]
+) -> list[int]:
+ result: list[int] = []
+ cursor = 0
+ for edit in edits:
+ if (
+ not 0 <= edit.start <= edit.stop <= len(rendered_prefix)
+ or edit.start < cursor
+ ):
+ raise ValueError("token-prefix edits must be ordered and non-overlapping")
+ result.extend(rendered_prefix[cursor : edit.start])
+ result.extend(edit.replacement)
+ cursor = edit.stop
+ result.extend(rendered_prefix[cursor:])
+ return result
+
+
+def prefix_edits(
+ rendered_prefix: Sequence[int], raw_prefix: Sequence[int]
+) -> tuple[PrefixEdit, ...]:
+ """Return an exact, deterministic edit script without nonlinear alignment."""
+ if rendered_prefix == raw_prefix:
+ return ()
+ if len(rendered_prefix) == len(raw_prefix):
+ edits: list[PrefixEdit] = []
+ start: int | None = None
+ for index, (rendered, raw) in enumerate(
+ zip(rendered_prefix, raw_prefix, strict=True)
+ ):
+ if rendered != raw and start is None:
+ start = index
+ elif rendered == raw and start is not None:
+ edits.append(PrefixEdit(start, index, tuple(raw_prefix[start:index])))
+ start = None
+ if start is not None:
+ edits.append(PrefixEdit(start, len(raw_prefix), tuple(raw_prefix[start:])))
+ return tuple(edits)
+
+ start = 0
+ limit = min(len(rendered_prefix), len(raw_prefix))
+ while start < limit and rendered_prefix[start] == raw_prefix[start]:
+ start += 1
+ suffix = 0
+ while (
+ suffix < len(rendered_prefix) - start
+ and suffix < len(raw_prefix) - start
+ and rendered_prefix[-1 - suffix] == raw_prefix[-1 - suffix]
+ ):
+ suffix += 1
+ raw_stop = len(raw_prefix) - suffix
+ return (
+ PrefixEdit(
+ start,
+ len(rendered_prefix) - suffix,
+ tuple(raw_prefix[start:raw_stop]),
+ ),
+ )
+
+
+def compact_prefix_candidate(
+ rendered_prefix: Sequence[int],
+ raw_prefix: Sequence[int],
+ edits: Sequence[PrefixEdit] | None = None,
+) -> CompactPrefixCandidate:
+ exact_edits = tuple(
+ prefix_edits(rendered_prefix, raw_prefix) if edits is None else edits
+ )
+ if apply_prefix_edits(rendered_prefix, exact_edits) != list(raw_prefix):
+ raise ValueError("token-prefix edits do not reconstruct the raw prefix")
+ return CompactPrefixCandidate(
+ rendered_length=len(rendered_prefix),
+ rendered_digest=token_digest(rendered_prefix),
+ raw_digest=token_digest(raw_prefix),
+ edits=exact_edits,
+ )
+
+
+def resolve_compact_prefix(
+ rendered: Sequence[int], candidates: Sequence[CompactPrefixCandidate]
+) -> PrefixMatch | None:
+ eligible = sorted(
+ (
+ candidate
+ for candidate in candidates
+ if 0 < candidate.rendered_length <= len(rendered)
+ ),
+ key=lambda candidate: candidate.rendered_length,
+ reverse=True,
+ )
+ if not eligible:
+ return None
+ digests = token_prefix_digests(
+ rendered, [candidate.rendered_length for candidate in eligible]
+ )
+ for candidate in eligible:
+ if digests[candidate.rendered_length] != candidate.rendered_digest:
+ continue
+ prefix = rendered[: candidate.rendered_length]
+ raw = apply_prefix_edits(prefix, candidate.edits)
+ if token_digest(raw) != candidate.raw_digest:
+ continue
+ return PrefixMatch(
+ rendered_length=candidate.rendered_length,
+ raw_prefix=tuple(raw),
+ edits=candidate.edits,
+ )
+ return None
+
+
+def compact_candidate_payload(
+ candidate: CompactPrefixCandidate,
+) -> dict[str, object]:
+ return {
+ "rendered_length": candidate.rendered_length,
+ "rendered_digest": candidate.rendered_digest,
+ "raw_digest": candidate.raw_digest,
+ "edits": [
+ [edit.start, edit.stop, list(edit.replacement)] for edit in candidate.edits
+ ],
+ }
+
+
+def compact_candidate_from_payload(value: object) -> CompactPrefixCandidate:
+ if not isinstance(value, dict):
+ raise ValueError("token-prefix candidate must be an object")
+ payload: dict[str, object] = {
+ key: item for key, item in value.items() if isinstance(key, str)
+ }
+ if len(payload) != len(value):
+ raise ValueError("token-prefix candidate keys must be strings")
+ rendered_length = payload.get("rendered_length")
+ rendered_digest = payload.get("rendered_digest")
+ raw_digest = payload.get("raw_digest")
+ raw_edits = payload.get("edits")
+ if (
+ not isinstance(rendered_length, int)
+ or isinstance(rendered_length, bool)
+ or not 0 < rendered_length <= _MAX_PREFIX_TOKENS
+ or not isinstance(rendered_digest, str)
+ or not isinstance(raw_digest, str)
+ or any(
+ len(digest) != 64
+ or any(character not in string.hexdigits for character in digest)
+ for digest in (rendered_digest, raw_digest)
+ )
+ or not isinstance(raw_edits, list)
+ or len(raw_edits) > 256
+ ):
+ raise ValueError("invalid token-prefix candidate")
+ edits: list[PrefixEdit] = []
+ cursor = 0
+ replacement_count = 0
+ for raw_edit in raw_edits:
+ if not isinstance(raw_edit, list) or len(raw_edit) != 3:
+ raise ValueError("invalid token-prefix edit")
+ start = raw_edit[0]
+ stop = raw_edit[1]
+ replacement = raw_edit[2]
+ if (
+ not isinstance(start, int)
+ or isinstance(start, bool)
+ or not isinstance(stop, int)
+ or isinstance(stop, bool)
+ or not isinstance(replacement, list)
+ ):
+ raise ValueError("invalid token-prefix edit")
+ replacement_ids: list[int] = []
+ for token_id in replacement:
+ if not isinstance(token_id, int) or isinstance(token_id, bool):
+ raise ValueError("invalid token-prefix edit")
+ replacement_ids.append(token_id)
+ if not cursor <= start <= stop <= rendered_length:
+ raise ValueError("token-prefix edits must be ordered and non-overlapping")
+ _token_bytes(replacement_ids)
+ replacement_count += len(replacement_ids)
+ if replacement_count > COMPACT_PREFIX_MAX_REPLACEMENT_IDS:
+ raise ValueError("token-prefix candidate has too many replacement tokens")
+ edits.append(PrefixEdit(start, stop, tuple(replacement_ids)))
+ cursor = stop
+ return CompactPrefixCandidate(
+ rendered_length=rendered_length,
+ rendered_digest=rendered_digest.lower(),
+ raw_digest=raw_digest.lower(),
+ edits=tuple(edits),
+ )
+
+
+def compact_candidates_header(
+ candidates: Sequence[CompactPrefixCandidate],
+) -> str | None:
+ """Encode a complete candidate snapshot, or decline unsafe header growth."""
+ prefix = f'{{"version":"{COMPACT_PREFIX_VERSION}","candidates":['
+ suffix = "]}"
+ parts = [prefix]
+ size = len(prefix) + len(suffix)
+ for candidate in candidates:
+ encoded = json.dumps(
+ compact_candidate_payload(candidate), separators=(",", ":")
+ )
+ separator = int(len(parts) > 1)
+ size += separator + len(encoded)
+ if size > COMPACT_PREFIX_HEADER_DECODE_LIMIT:
+ return None
+ if separator:
+ parts.append(",")
+ parts.append(encoded)
+ parts.append(suffix)
+ encoded = base64.urlsafe_b64encode(
+ gzip.compress("".join(parts).encode(), compresslevel=1)
+ ).decode()
+ return encoded if len(encoded) <= COMPACT_PREFIX_HEADER_LIMIT else None
+
+
+def compact_candidates_from_header(
+ value: str,
+) -> tuple[CompactPrefixCandidate, ...]:
+ if len(value) > COMPACT_PREFIX_HEADER_LIMIT:
+ raise ValueError("token-prefix candidate header is too large")
+ try:
+ compressed = base64.b64decode(value, altchars=b"-_", validate=True)
+ decompressor = zlib.decompressobj(16 + zlib.MAX_WBITS)
+ decoded = decompressor.decompress(
+ compressed, COMPACT_PREFIX_HEADER_DECODE_LIMIT + 1
+ )
+ if (
+ len(decoded) > COMPACT_PREFIX_HEADER_DECODE_LIMIT
+ or decompressor.unconsumed_tail
+ ):
+ raise ValueError("token-prefix candidate header expands beyond its limit")
+ decoded += decompressor.flush(
+ COMPACT_PREFIX_HEADER_DECODE_LIMIT + 1 - len(decoded)
+ )
+ if not decompressor.eof or decompressor.unused_data:
+ raise ValueError("token-prefix candidate header is malformed")
+ payload = json.loads(decoded)
+ except (
+ binascii.Error,
+ json.JSONDecodeError,
+ UnicodeDecodeError,
+ zlib.error,
+ ) as error:
+ raise ValueError("token-prefix candidate header is malformed") from error
+ if (
+ not isinstance(payload, dict)
+ or payload.get("version") != COMPACT_PREFIX_VERSION
+ or not isinstance(payload.get("candidates"), list)
+ or len(payload["candidates"]) > COMPACT_PREFIX_MAX_CANDIDATES
+ ):
+ raise ValueError("token-prefix candidate header is malformed")
+ return tuple(
+ compact_candidate_from_payload(candidate) for candidate in payload["candidates"]
+ )
+
+
+class _Edge:
+ __slots__ = ("label", "child")
+
+ def __init__(self, label: tuple[int, ...], child: _Node) -> None:
+ self.label = label
+ self.child = child
+
+
+@dataclass(slots=True)
+class _Variant:
+ lineages: OrderedDict[str, int]
+ edits: tuple[PrefixEdit, ...]
+
+
+class _Node:
+ __slots__ = ("children", "parent", "parent_token", "rendered_length", "variants")
+
+ def __init__(
+ self,
+ parent: _Node | None = None,
+ parent_token: int | None = None,
+ rendered_length: int = 0,
+ ) -> None:
+ self.children: dict[int, _Edge] = {}
+ self.parent = parent
+ self.parent_token = parent_token
+ self.rendered_length = rendered_length
+ self.variants: dict[tuple[int, ...], _Variant] = {}
+
+
+def _common_prefix_length(
+ values: Sequence[int], start: int, label: tuple[int, ...]
+) -> int:
+ limit = min(len(values) - start, len(label))
+ index = 0
+ while index < limit and values[start + index] == label[index]:
+ index += 1
+ return index
+
+
+class TokenPrefixCache:
+ """A bounded radix trie retaining distinct observed raw tokenizations."""
+
+ def __init__(
+ self,
+ *,
+ max_variants: int = 4096,
+ max_token_ids: int = 2_000_000,
+ max_lineages_per_variant: int = 64,
+ ) -> None:
+ if min(max_variants, max_token_ids, max_lineages_per_variant) <= 0:
+ raise ValueError("token-prefix cache bounds must be positive")
+ self._root = _Node()
+ self._lru: OrderedDict[_Node, int] = OrderedDict()
+ self._variants = 0
+ self._max_variants = max_variants
+ self._max_token_ids = max_token_ids
+ self._max_lineages_per_variant = max_lineages_per_variant
+ self._token_ids = 0
+ self._observation = 0
+
+ def _matching_nodes(self, rendered_tokens: Sequence[int]) -> list[_Node]:
+ node = self._root
+ index = 0
+ candidates: list[_Node] = []
+ while index < len(rendered_tokens):
+ edge = node.children.get(rendered_tokens[index])
+ if edge is None:
+ break
+ matched = _common_prefix_length(rendered_tokens, index, edge.label)
+ if matched != len(edge.label):
+ break
+ index += matched
+ node = edge.child
+ if node.variants:
+ candidates.append(node)
+ return candidates
+
+ def lookup(
+ self, rendered_tokens: Sequence[int], lineage: str | None
+ ) -> PrefixMatch | None:
+ for candidate in reversed(self._matching_nodes(rendered_tokens)):
+ variants = candidate.variants
+ # Without a rollout identifier, never choose between distinct raw
+ # tokenizations of the same rendered history.
+ if lineage is None:
+ if len(variants) == 1:
+ raw, variant = next(iter(variants.items()))
+ self._lru.move_to_end(candidate)
+ return PrefixMatch(candidate.rendered_length, raw, variant.edits)
+ continue
+ matching = [
+ (variant.lineages[lineage], raw, variant.edits)
+ for raw, variant in variants.items()
+ if lineage in variant.lineages
+ ]
+ if not matching:
+ continue
+ _, raw_prefix, edits = max(matching, key=lambda value: value[0])
+ self._lru.move_to_end(candidate)
+ lineages = variants[raw_prefix].lineages
+ lineages.move_to_end(lineage)
+ return PrefixMatch(candidate.rendered_length, raw_prefix, edits)
+ return None
+
+ def insert(
+ self,
+ rendered_prefix: Sequence[int],
+ raw_prefix: Sequence[int],
+ lineage: str,
+ edits: Sequence[PrefixEdit] | None = None,
+ ) -> bool:
+ if not rendered_prefix or not raw_prefix or not lineage:
+ return False
+ raw = tuple(raw_prefix)
+ cost = len(rendered_prefix) + len(raw)
+ if cost > self._max_token_ids:
+ return False
+ node = self._root
+ index = 0
+ while index < len(rendered_prefix):
+ token = rendered_prefix[index]
+ edge = node.children.get(token)
+ if edge is None:
+ child = _Node(
+ parent=node,
+ parent_token=token,
+ rendered_length=len(rendered_prefix),
+ )
+ node.children[token] = _Edge(tuple(rendered_prefix[index:]), child)
+ node = child
+ break
+
+ matched = _common_prefix_length(rendered_prefix, index, edge.label)
+ if matched == len(edge.label):
+ index += matched
+ node = edge.child
+ continue
+
+ middle = _Node(
+ parent=node,
+ parent_token=token,
+ rendered_length=index + matched,
+ )
+ old_suffix = edge.label[matched:]
+ old_child = edge.child
+ old_child.parent = middle
+ old_child.parent_token = old_suffix[0]
+ middle.children[old_suffix[0]] = _Edge(old_suffix, old_child)
+ edge.label = edge.label[:matched]
+ edge.child = middle
+ node = middle
+ index += matched
+ if index < len(rendered_prefix):
+ new_token = rendered_prefix[index]
+ child = _Node(
+ parent=node,
+ parent_token=new_token,
+ rendered_length=len(rendered_prefix),
+ )
+ node.children[new_token] = _Edge(tuple(rendered_prefix[index:]), child)
+ node = child
+ break
+
+ exact_edits = tuple(
+ prefix_edits(rendered_prefix, raw_prefix) if edits is None else edits
+ )
+ if apply_prefix_edits(rendered_prefix, exact_edits) != list(raw_prefix):
+ raise ValueError("token-prefix edits do not reconstruct the raw prefix")
+ variant = node.variants.get(raw)
+ if variant is None:
+ variant = node.variants[raw] = _Variant(OrderedDict(), exact_edits)
+ self._lru[node] = self._lru.get(node, 0) + cost
+ self._variants += 1
+ self._token_ids += cost
+ else:
+ variant.edits = exact_edits
+ self._lru.move_to_end(node)
+ lineages = variant.lineages
+ self._observation += 1
+ lineages[lineage] = self._observation
+ lineages.move_to_end(lineage)
+ while len(lineages) > self._max_lineages_per_variant:
+ lineages.popitem(last=False)
+ if node in self._lru:
+ self._lru.move_to_end(node)
+ self._evict()
+ return True
+
+ def _evict(self) -> None:
+ while self._lru and (
+ self._variants > self._max_variants or self._token_ids > self._max_token_ids
+ ):
+ self.evict_oldest()
+
+ @property
+ def token_ids(self) -> int:
+ return self._token_ids
+
+ def evict_oldest(self) -> bool:
+ if not self._lru:
+ return False
+ node, cost = self._lru.popitem(last=False)
+ self._token_ids -= cost
+ self._variants -= len(node.variants)
+ node.variants.clear()
+ self._prune(node)
+ return True
+
+ def _prune(self, node: _Node) -> None:
+ while node.parent is not None:
+ parent = node.parent
+ token = node.parent_token
+ assert token is not None
+ if not node.variants and not node.children:
+ del parent.children[token]
+ node = parent
+ continue
+ if not node.variants and len(node.children) == 1:
+ child_edge = next(iter(node.children.values()))
+ parent_edge = parent.children[token]
+ parent_edge.label += child_edge.label
+ parent_edge.child = child_edge.child
+ child_edge.child.parent = parent
+ child_edge.child.parent_token = token
+ node = parent
+ continue
+ break
+
+
+class TokenPrefixStore:
+ """LRU-bounded collection of model/tokenizer-scoped prefix tries."""
+
+ def __init__(
+ self,
+ *,
+ max_scopes: int = 64,
+ max_variants_per_scope: int = 4096,
+ max_token_ids_per_scope: int = 2_000_000,
+ max_token_ids: int = 4_000_000,
+ ) -> None:
+ if min(max_scopes, max_token_ids) <= 0:
+ raise ValueError("token-prefix store bounds must be positive")
+ self._max_scopes = max_scopes
+ self._max_variants = max_variants_per_scope
+ self._max_token_ids = max_token_ids_per_scope
+ self._max_total_token_ids = max_token_ids
+ self._token_ids = 0
+ self._scopes: OrderedDict[str, TokenPrefixCache] = OrderedDict()
+
+ def _cache(self, scope: str) -> TokenPrefixCache:
+ cache = self._scopes.get(scope)
+ if cache is None:
+ cache = TokenPrefixCache(
+ max_variants=self._max_variants,
+ max_token_ids=self._max_token_ids,
+ )
+ self._scopes[scope] = cache
+ while len(self._scopes) > self._max_scopes:
+ _, removed = self._scopes.popitem(last=False)
+ self._token_ids -= removed.token_ids
+ else:
+ self._scopes.move_to_end(scope)
+ return cache
+
+ def lookup(
+ self, scope: str, rendered_tokens: Sequence[int], lineage: str | None
+ ) -> PrefixMatch | None:
+ cache = self._scopes.get(scope)
+ if cache is None:
+ return None
+ self._scopes.move_to_end(scope)
+ return cache.lookup(rendered_tokens, lineage)
+
+ def insert(
+ self,
+ scope: str,
+ rendered_prefix: Sequence[int],
+ raw_prefix: Sequence[int],
+ lineage: str,
+ edits: Sequence[PrefixEdit] | None = None,
+ ) -> bool:
+ cache = self._cache(scope)
+ before = cache.token_ids
+ inserted = cache.insert(rendered_prefix, raw_prefix, lineage, edits)
+ self._token_ids += cache.token_ids - before
+ while self._token_ids > self._max_total_token_ids and self._scopes:
+ oldest_scope, oldest = next(iter(self._scopes.items()))
+ before = oldest.token_ids
+ if not oldest.evict_oldest():
+ self._scopes.pop(oldest_scope)
+ continue
+ self._token_ids -= before - oldest.token_ids
+ if oldest.token_ids == 0:
+ self._scopes.pop(oldest_scope)
+ return inserted
+
+
+class CompactPrefixStore:
+ """Bounded shared snapshots keyed by scope and logical rollout lineage."""
+
+ def __init__(
+ self,
+ *,
+ max_candidates: int = 262_144,
+ max_candidates_per_lineage: int = COMPACT_PREFIX_MAX_CANDIDATES,
+ max_edit_token_ids: int = 1_000_000,
+ max_serialized_bytes_per_lineage: int = COMPACT_PREFIX_RESPONSE_LIMIT - 1024,
+ max_staged_attempts: int = 262_144,
+ staged_ttl: float = 300,
+ ) -> None:
+ if (
+ min(
+ max_candidates,
+ max_candidates_per_lineage,
+ max_edit_token_ids,
+ max_serialized_bytes_per_lineage,
+ max_staged_attempts,
+ staged_ttl,
+ )
+ <= 0
+ ):
+ raise ValueError("compact token-prefix store bounds must be positive")
+ self._max_candidates = max_candidates
+ self._max_per_lineage = max_candidates_per_lineage
+ self._max_edit_token_ids = max_edit_token_ids
+ self._max_serialized_bytes_per_lineage = max_serialized_bytes_per_lineage
+ self._edit_token_ids = 0
+ self._entries: OrderedDict[
+ tuple[str, str, int, str], tuple[CompactPrefixCandidate, int, int]
+ ] = OrderedDict()
+ self._lineages: dict[
+ tuple[str, str], OrderedDict[tuple[str, str, int, str], None]
+ ] = {}
+ self._lineage_bytes: dict[tuple[str, str], int] = {}
+ self._max_staged_attempts = max_staged_attempts
+ self._staged_ttl = staged_ttl
+ self._staged: OrderedDict[
+ str,
+ tuple[
+ float,
+ list[tuple[str, str, CompactPrefixCandidate]],
+ int,
+ ],
+ ] = OrderedDict()
+ self._staged_candidates = 0
+ self._staged_edit_token_ids = 0
+ self._staged_evictions = 0
+ self._missing_commits = 0
+
+ @property
+ def staged_evictions(self) -> int:
+ return self._staged_evictions
+
+ @property
+ def missing_commits(self) -> int:
+ return self._missing_commits
+
+ def insert(
+ self,
+ scope: str,
+ lineage: str,
+ candidate: CompactPrefixCandidate,
+ ) -> bool:
+ cost = sum(len(edit.replacement) for edit in candidate.edits)
+ serialized = (
+ len(
+ json.dumps(
+ compact_candidate_payload(candidate), separators=(",", ":")
+ ).encode()
+ )
+ + 1
+ )
+ if (
+ cost > self._max_edit_token_ids
+ or serialized > self._max_serialized_bytes_per_lineage
+ ):
+ return False
+ key = (
+ scope,
+ lineage,
+ candidate.rendered_length,
+ candidate.rendered_digest,
+ )
+ previous = self._entries.pop(key, None)
+ if previous is not None:
+ self._edit_token_ids -= previous[1]
+ self._lineage_bytes[(scope, lineage)] -= previous[2]
+ self._entries[key] = (candidate, cost, serialized)
+ self._edit_token_ids += cost
+ lineage_key = (scope, lineage)
+ self._lineage_bytes[lineage_key] = (
+ self._lineage_bytes.get(lineage_key, 0) + serialized
+ )
+ entries = self._lineages.setdefault(lineage_key, OrderedDict())
+ entries.pop(key, None)
+ entries[key] = None
+ while (
+ len(entries) > self._max_per_lineage
+ or self._lineage_bytes[lineage_key] > self._max_serialized_bytes_per_lineage
+ ):
+ self._remove(next(iter(entries)))
+ while (
+ len(self._entries) > self._max_candidates
+ or self._edit_token_ids > self._max_edit_token_ids
+ ):
+ self._remove(next(iter(self._entries)))
+ return key in self._entries
+
+ def candidates(
+ self, scope: str, lineage: str, max_rendered_length: int
+ ) -> list[CompactPrefixCandidate]:
+ entries = self._lineages.get((scope, lineage))
+ if entries is None:
+ return []
+ ordered = [
+ self._entries[key][0]
+ for key in reversed(entries)
+ if key in self._entries
+ and self._entries[key][0].rendered_length <= max_rendered_length
+ ]
+ # Python's sort is stable, so equal-length variants remain newest first.
+ ordered.sort(key=lambda candidate: candidate.rendered_length, reverse=True)
+ return ordered
+
+ def stage(
+ self,
+ attempt: str,
+ observations: Sequence[tuple[str, str, CompactPrefixCandidate]],
+ ) -> bool:
+ """Retain observations invisibly until their upstream attempt commits."""
+ now = time.monotonic()
+ self._expire_staged(now)
+ existing = self._staged.pop(attempt, None)
+ entries = [] if existing is None else existing[1]
+ cost = 0 if existing is None else existing[2]
+ if existing is not None:
+ self._staged_candidates -= len(entries)
+ self._staged_edit_token_ids -= cost
+ combined = OrderedDict(
+ (
+ (
+ scope,
+ lineage,
+ candidate.rendered_length,
+ candidate.rendered_digest,
+ ),
+ (scope, lineage, candidate),
+ )
+ for scope, lineage, candidate in entries
+ )
+ for scope, lineage, candidate in observations:
+ candidate_cost = sum(len(edit.replacement) for edit in candidate.edits)
+ if candidate_cost <= self._max_edit_token_ids:
+ key = (
+ scope,
+ lineage,
+ candidate.rendered_length,
+ candidate.rendered_digest,
+ )
+ combined.pop(key, None)
+ combined[key] = (scope, lineage, candidate)
+ entries = list(combined.values())
+ cost = sum(
+ len(edit.replacement)
+ for _, _, candidate in entries
+ for edit in candidate.edits
+ )
+ if entries:
+ self._staged[attempt] = (now, entries, cost)
+ self._staged_candidates += len(entries)
+ self._staged_edit_token_ids += cost
+ while self._staged and (
+ len(self._staged) > self._max_staged_attempts
+ or self._staged_candidates > self._max_candidates
+ or self._staged_edit_token_ids > self._max_edit_token_ids
+ ):
+ self._remove_staged(next(iter(self._staged)))
+ self._staged_evictions += 1
+ return attempt in self._staged
+
+ def commit(self, attempt: str) -> bool:
+ """Publish only observations produced by the selected upstream attempt."""
+ self._expire_staged(time.monotonic())
+ staged = self._staged.get(attempt)
+ if staged is None:
+ self._missing_commits += 1
+ return False
+ observations = list(staged[1])
+ self._remove_staged(attempt)
+ inserted = [
+ self.insert(scope, lineage, candidate)
+ for scope, lineage, candidate in observations
+ ]
+ return all(inserted)
+
+ def _expire_staged(self, now: float) -> None:
+ while self._staged:
+ attempt, (created, _, _) = next(iter(self._staged.items()))
+ if now - created < self._staged_ttl:
+ break
+ self._remove_staged(attempt)
+
+ def _remove_staged(self, attempt: str) -> None:
+ removed = self._staged.pop(attempt, None)
+ if removed is None:
+ return
+ self._staged_candidates -= len(removed[1])
+ self._staged_edit_token_ids -= removed[2]
+
+ def _remove(self, key: tuple[str, str, int, str]) -> None:
+ removed = self._entries.pop(key, None)
+ if removed is None:
+ return
+ self._edit_token_ids -= removed[1]
+ lineage_key = key[:2]
+ self._lineage_bytes[lineage_key] -= removed[2]
+ entries = self._lineages[lineage_key]
+ entries.pop(key, None)
+ if not entries:
+ del self._lineages[lineage_key]
+ del self._lineage_bytes[lineage_key]
+
+
+class TokenPrefixRuntime:
+ """Local hot cache with shared lookup and batched replication."""
+
+ def __init__(
+ self,
+ *,
+ lookup_shared: Callable[[str, list[int], str], Awaitable[PrefixMatch | None]],
+ insert_shared_many: Callable[[list[_SharedPrefixEntry]], Awaitable[None]],
+ max_pending_batches: int = 32,
+ ) -> None:
+ if max_pending_batches <= 0:
+ raise ValueError("replication bound must be positive")
+ self._local = TokenPrefixStore()
+ self._lookup_shared = lookup_shared
+ self._insert_shared_many = insert_shared_many
+ self._pending: asyncio.Queue[list[_SharedPrefixEntry]] = asyncio.Queue(
+ max_pending_batches
+ )
+ self._replicator: asyncio.Task[None] | None = None
+ self._closed = False
+
+ async def rewrite_with_edits(
+ self,
+ scope: str,
+ rendered_tokens: Sequence[int],
+ lineage: str,
+ *,
+ fallback_lineage: str | None = None,
+ shared_candidate: bool,
+ shared_match: PrefixMatch | None = None,
+ shared_resolved: bool = False,
+ ) -> tuple[list[int], list[int], tuple[PrefixEdit, ...]]:
+ """Return canonical/input tokens plus their exact sparse substitution."""
+ canonical = list(rendered_tokens)
+ local_match = self._local.lookup(scope, canonical, lineage)
+ if local_match is None and fallback_lineage is not None:
+ local_match = self._local.lookup(scope, canonical, fallback_lineage)
+ if local_match is not None:
+ self._local.insert(
+ scope,
+ canonical[: local_match.rendered_length],
+ local_match.raw_prefix,
+ lineage,
+ local_match.edits,
+ )
+ if shared_candidate and not shared_resolved:
+ shared_match = await self._lookup_shared(scope, canonical, lineage)
+ if shared_match is None and fallback_lineage is not None:
+ shared_match = await self._lookup_shared(
+ scope, canonical, fallback_lineage
+ )
+ if shared_match is not None and (
+ local_match is None
+ or shared_match.rendered_length > local_match.rendered_length
+ ):
+ exact_edits = shared_match.edits
+ rendered_prefix = canonical[: shared_match.rendered_length]
+ if not exact_edits and list(shared_match.raw_prefix) != rendered_prefix:
+ exact_edits = prefix_edits(rendered_prefix, shared_match.raw_prefix)
+ self._local.insert(
+ scope,
+ rendered_prefix,
+ shared_match.raw_prefix,
+ lineage,
+ exact_edits,
+ )
+ match = PrefixMatch(
+ shared_match.rendered_length,
+ shared_match.raw_prefix,
+ exact_edits,
+ )
+ else:
+ match = local_match
+ if match is None:
+ return canonical, canonical.copy(), ()
+ return (
+ canonical,
+ [*match.raw_prefix, *canonical[match.rendered_length :]],
+ match.edits,
+ )
+
+ async def insert_many(
+ self,
+ entries: Sequence[_SharedPrefixEntry],
+ ) -> None:
+ normalized: list[_SharedPrefixEntry] = []
+ for scope, rendered, raw, lineage, edits in entries:
+ exact_edits = tuple(edits)
+ self._local.insert(scope, rendered, raw, lineage, exact_edits)
+ normalized.append((scope, rendered, raw, lineage, exact_edits))
+ if not normalized or self._closed:
+ return
+ try:
+ self._pending.put_nowait(normalized)
+ except asyncio.QueueFull:
+ # Let the already-scheduled replicator drain a simultaneous burst
+ # once before dropping cross-replica work.
+ self._ensure_replicator()
+ await asyncio.sleep(0)
+ if self._closed:
+ return
+ try:
+ self._pending.put_nowait(normalized)
+ except asyncio.QueueFull:
+ logger.warning(
+ "shared token-prefix replication queue is full; "
+ "retaining local mappings only"
+ )
+ return
+ self._ensure_replicator()
+
+ def _ensure_replicator(self) -> None:
+ if self._replicator is None or self._replicator.done():
+ self._replicator = asyncio.create_task(
+ self._replicate(), name="caladan-token-prefix-replication"
+ )
+
+ async def _replicate(self) -> None:
+ try:
+ while not self._pending.empty():
+ batches = [self._pending.get_nowait()]
+ while True:
+ try:
+ batches.append(self._pending.get_nowait())
+ except asyncio.QueueEmpty:
+ break
+ try:
+ entries = [entry for batch in batches for entry in batch]
+ await self._insert_shared_many(entries)
+ except asyncio.CancelledError:
+ raise
+ except Exception:
+ logger.warning(
+ "shared token-prefix insertion failed; "
+ "retaining local mappings",
+ exc_info=True,
+ )
+ finally:
+ for _ in batches:
+ self._pending.task_done()
+ finally:
+ if self._replicator is asyncio.current_task():
+ self._replicator = None
+ if not self._closed and not self._pending.empty():
+ self._ensure_replicator()
+
+ async def close(self, *, timeout: float = 2.0) -> bool:
+ """Drain pending replication for at most timeout seconds, then stop."""
+ self._closed = True
+ drained = await self.flush(timeout=timeout)
+ replicator = self._replicator
+ if replicator is not None:
+ replicator.cancel()
+ with suppress(asyncio.CancelledError):
+ await replicator
+ if self._replicator is replicator:
+ self._replicator = None
+ while not self._pending.empty():
+ self._pending.get_nowait()
+ self._pending.task_done()
+ return drained
+
+ async def flush(self, *, timeout: float | None = None) -> bool:
+ """Wait until queued shared replication has completed."""
+ if self._replicator is None:
+ return True
+ try:
+ await asyncio.wait_for(self._pending.join(), timeout)
+ except TimeoutError:
+ logger.warning("token-prefix replication did not drain before timeout")
+ return False
+ return True
diff --git a/src/art_inference/vllm.py b/src/art_inference/vllm.py
new file mode 100644
index 000000000..ec007a576
--- /dev/null
+++ b/src/art_inference/vllm.py
@@ -0,0 +1,445 @@
+"""Preserve sampled history before an ART-owned vLLM engine serves the next turn."""
+
+from __future__ import annotations
+
+from collections.abc import AsyncIterator, Awaitable, Callable
+from contextvars import ContextVar
+import copy
+from dataclasses import dataclass, field
+from functools import wraps
+import hashlib
+import importlib
+import inspect
+import json
+from types import MethodType
+from typing import Any
+
+from .append_only import (
+ aligned_values,
+ chat_response_prefixes,
+ merge_chat_delta,
+ output_prefix_observations,
+ patch_deepseek_renderer,
+ patch_harmony,
+ preserves_history,
+ shifted_span,
+)
+from .chat_template import (
+ chat_template_with_preserved_thinking,
+ configure_preserved_thinking_chat_template,
+)
+from .token_prefix import TokenPrefixStore, prefix_edits
+
+
+@dataclass
+class _Turn:
+ scope: str
+ tokenizer: Any
+ rendered: list[int] = field(default_factory=list)
+ prompt: list[int] = field(default_factory=list)
+ outputs: dict[int, tuple[list[int], bool]] = field(default_factory=dict)
+ external_request: Any = None
+ external_observer: Callable[[Any, list[Any]], Awaitable[None]] | None = None
+ enabled: bool = True
+
+
+_CURRENT: ContextVar[_Turn | None] = ContextVar("art_history", default=None)
+_PREFIXES = TokenPrefixStore()
+_DROP_THINKING: ContextVar[bool] = ContextVar("art_vllm_drop_thinking", default=False)
+
+
+def _configure_native_tokenizer(tokenizer, importer):
+ apply = getattr(tokenizer, "apply_chat_template", None)
+ if (
+ not callable(apply)
+ or getattr(apply, "__module__", None) != "vllm.tokenizers.deepseek_v32"
+ ):
+ return
+ module = importer("vllm.tokenizers.deepseek_v32")
+ if not getattr(module, "_art_append_only", False):
+ encode = module.encode_messages
+ patch_deepseek_renderer(
+ importer("vllm.tokenizers.deepseek_v32_encoding"),
+ lambda: not _DROP_THINKING.get(),
+ prefix_only=True,
+ )
+
+ @wraps(encode)
+ def encode_messages(*args, **kwargs):
+ kwargs["drop_thinking"] = _DROP_THINKING.get()
+ return encode(*args, **kwargs)
+
+ module.encode_messages = encode_messages
+ module._art_append_only = True
+
+ @wraps(apply)
+ def apply_template(self, *args, **kwargs):
+ token = _DROP_THINKING.set(not preserves_history(kwargs))
+ try:
+ return apply(*args, **kwargs)
+ finally:
+ _DROP_THINKING.reset(token)
+
+ tokenizer.apply_chat_template = MethodType(apply_template, tokenizer)
+
+
+def set_external_history_observer(
+ observer: Callable[[Any, list[Any]], Awaitable[None]],
+ importer=importlib.import_module,
+) -> None:
+ """Use a hosting service's transactional/distributed store for its requests."""
+ base = importer("vllm.renderers.base").BaseRenderer
+ base._art_history_external_observer = staticmethod(observer)
+
+
+async def _observe(turn: _Turn, entries) -> None:
+ if turn.external_request is not None:
+ assert turn.external_observer is not None
+ await turn.external_observer(turn.external_request, entries)
+ else:
+ for rendered, raw, edits in entries:
+ _PREFIXES.insert(turn.scope, rendered, raw, "content", edits)
+
+
+def _tokens(prompt: Any) -> list[int] | None:
+ if isinstance(prompt, dict):
+ tokens = prompt.get("prompt_token_ids")
+ if isinstance(tokens, (list, tuple)):
+ return list(tokens)
+ return None
+
+
+def replace_prompt_tokens(prompt, tokens, edits=None):
+ """Keep native multimodal offsets and per-token metadata aligned."""
+ result = {**prompt, "prompt_token_ids": tokens}
+ if tokens == prompt["prompt_token_ids"]:
+ return result
+ edits = prefix_edits(prompt["prompt_token_ids"], tokens) if edits is None else edits
+ if prompt.get("mm_placeholders"):
+ result["mm_placeholders"] = {}
+ for modality, spans in prompt["mm_placeholders"].items():
+ result["mm_placeholders"][modality] = []
+ for span in spans:
+ shifted = copy.copy(span)
+ shifted.offset = shifted_span(span.offset, span.length, edits)
+ result["mm_placeholders"][modality].append(shifted)
+ for name in (
+ "is_token_ids",
+ "prompt_is_token_ids",
+ "token_type_ids",
+ "assistant_tokens_mask",
+ ):
+ if prompt.get(name) is not None:
+ fill = (
+ True
+ if name in ("is_token_ids", "prompt_is_token_ids")
+ else (0 if name == "token_type_ids" else None)
+ )
+ result[name] = aligned_values(prompt[name], edits, fill=fill)
+ # Text offsets refer to the normalized string, which is no longer the
+ # source of these exact IDs. They are optional engine metadata.
+ result.pop("prompt_token_offsets", None)
+ return result
+
+
+def patch_history(importer=importlib.import_module) -> None:
+ base = importer("vllm.renderers.base").BaseRenderer
+ if getattr(base, "_art_append_only", False):
+ return
+ init = base.__init__
+
+ @wraps(init)
+ def initialize(self, config, tokenizer):
+ if tokenizer is not None:
+ configure_preserved_thinking_chat_template(tokenizer)
+ _configure_native_tokenizer(tokenizer, importer)
+ init(self, config, tokenizer)
+
+ base.__init__ = initialize
+ base._art_append_only = True
+ patch_harmony(
+ importer("vllm.entrypoints.openai.parser.harmony_utils"),
+ lambda: (turn := _CURRENT.get()) is None or turn.enabled,
+ "vllm.",
+ )
+
+ params = importer("vllm.renderers.params").ChatParams
+ apply_kwargs = params.get_apply_chat_template_kwargs
+
+ @wraps(apply_kwargs)
+ def template_kwargs(self):
+ kwargs = {
+ "preserve_thinking": True,
+ "clear_thinking": False,
+ "drop_thinking": False,
+ **apply_kwargs(self),
+ }
+ if kwargs.get("chat_template") is not None:
+ kwargs["chat_template"] = chat_template_with_preserved_thinking(
+ kwargs["chat_template"]
+ )
+ return kwargs
+
+ params.get_apply_chat_template_kwargs = template_kwargs
+
+ responses_utils = importer("vllm.entrypoints.openai.responses.utils")
+ responses_serving = importer("vllm.entrypoints.openai.responses.serving")
+ construct = responses_utils.construct_input_messages
+
+ @wraps(construct)
+ def input_messages(*, prev_response_output=None, request_input, **kwargs):
+ # vLLM's previous_response_id path otherwise drops reasoning and tool
+ # calls, although its ordinary input path already knows how to retain
+ # both. Keep the protocol's existing instructions replacement semantics.
+ if prev_response_output:
+ request_input = [
+ *prev_response_output,
+ *(
+ [{"role": "user", "content": request_input}]
+ if isinstance(request_input, str)
+ else request_input
+ ),
+ ]
+ return construct(request_input=request_input, **kwargs)
+
+ setattr(responses_utils, "construct_input_messages", input_messages)
+ setattr(responses_serving, "construct_input_messages", input_messages)
+
+ engine = importer("vllm.v1.engine.async_llm").AsyncLLM
+ generate = engine.generate
+ signature = inspect.signature(generate)
+
+ @wraps(generate)
+ def serve(self, prompt, *args, **kwargs):
+ turn = _CURRENT.get()
+ tokens = _tokens(prompt)
+ if turn is None or not turn.enabled or tokens is None:
+ return generate(self, prompt, *args, **kwargs)
+ turn.rendered = tokens
+ # Batched completions run several generators concurrently in one API
+ # request. Each generator must publish only its own prompt and outputs.
+ outputs = {}
+ turn.outputs = outputs
+ sampling = signature.bind(self, prompt, *args, **kwargs).arguments.get(
+ "sampling_params"
+ )
+ delta = getattr(getattr(sampling, "output_kind", None), "name", None) == "DELTA"
+ match = (
+ _PREFIXES.lookup(turn.scope, tokens, None)
+ if turn.external_request is None
+ else None
+ )
+ raw_prompt = (
+ list(match.raw_prefix) + tokens[match.rendered_length :]
+ if match
+ else tokens
+ )
+ turn.prompt = raw_prompt
+ prompt = replace_prompt_tokens(prompt, raw_prompt, match.edits if match else ())
+
+ async def stream():
+ async for output in generate(self, prompt, *args, **kwargs):
+ for choice in output.outputs:
+ generated_ids = list(choice.token_ids)
+ if delta:
+ generated_ids = (
+ outputs.get(choice.index, ([], False))[0] + generated_ids
+ )
+ outputs[choice.index] = (
+ generated_ids,
+ choice.finish_reason not in (None, "length", "abort"),
+ )
+ yield output
+ # Decode/re-encode differences matter for completions and for chat
+ # renderers that already preserve the sampled text exactly.
+ for raw, _ in outputs.values():
+ if turn.external_request is None:
+ await _observe(
+ turn,
+ output_prefix_observations(
+ turn.tokenizer, tokens, raw_prompt, raw
+ ),
+ )
+
+ return stream()
+
+ engine.generate = serve
+ chat = importer("vllm.entrypoints.openai.chat_completion.serving").OpenAIServingChat
+ completion = importer(
+ "vllm.entrypoints.openai.completion.serving"
+ ).OpenAIServingCompletion
+
+ def install(owner, method, *, observe_chat=False, observe_responses=False):
+ original = getattr(owner, method)
+
+ @wraps(original)
+ async def create(self, request, raw_request=None):
+ headers = getattr(raw_request, "headers", {}) or {}
+ # Caladan supplies a distributed store and its own engine adapter.
+ # Its observer uses the same shared ART implementation.
+ external = bool(headers.get("x-caladan-prefix-scope"))
+ external_observer = getattr(base, "_art_history_external_observer", None)
+ if external and external_observer is None:
+ return await original(self, request, raw_request)
+ tokenizer = getattr(self.renderer, "tokenizer", None)
+ if tokenizer is None:
+ return await original(self, request, raw_request)
+ material = [
+ request.model,
+ getattr(request, "chat_template", None),
+ getattr(request, "chat_template_kwargs", None),
+ headers.get("authorization", ""),
+ ]
+ scope = hashlib.sha256(
+ json.dumps(material, sort_keys=True).encode()
+ ).hexdigest()
+ turn = _Turn(
+ scope,
+ tokenizer,
+ external_request=raw_request if external else None,
+ external_observer=external_observer,
+ enabled=preserves_history(
+ getattr(request, "chat_template_kwargs", None)
+ ),
+ )
+
+ async def observe(messages):
+ if not turn.enabled or not observe_chat or not turn.prompt:
+ return
+
+ async def render(value):
+ result = await self.render_chat_request(value)
+ if not isinstance(result, tuple) or len(result[1]) != 1:
+ return None
+ return _tokens(result[1][0])
+
+ choices = [
+ (message, *turn.outputs[index])
+ for index, message in messages.items()
+ if index in turn.outputs
+ ]
+ entries = await chat_response_prefixes(
+ tokenizer, request, turn.prompt, choices, render
+ )
+ await _observe(turn, entries)
+
+ async def observe_response(response):
+ if (
+ not observe_responses
+ or not turn.enabled
+ or getattr(self, "use_harmony", False)
+ or not getattr(response, "output", None)
+ or not turn.prompt
+ ):
+ return
+ messages = responses_utils.construct_chat_messages_with_tool_call(
+ response.output
+ )
+ # Built-in tool execution can contain several generations. The
+ # engine-level observations cover those individual turns.
+ if len(messages) != 1 or messages[0].get("role") != "assistant":
+ return
+ previous = None
+ if request.previous_response_id:
+ async with self.response_store_lock:
+ previous = self.response_store.get(request.previous_response_id)
+ conversation, inputs = await self._make_request(request, previous)
+ if len(inputs) != 1:
+ return
+ protocol = importer("vllm.entrypoints.openai.chat_completion.protocol")
+ tools = responses_utils.construct_tool_dicts(
+ request.tools, request.tool_choice
+ )
+ view = protocol.ChatCompletionRequest(
+ model=request.model,
+ messages=conversation,
+ chat_template_kwargs=self._effective_chat_template_kwargs(request),
+ )
+
+ async def render(value):
+ online = getattr(self, "online_renderer", None)
+ if online is None: # vLLM 0.24
+ online = self.openai_serving_render
+ _, inputs = await online.preprocess_chat(
+ value,
+ value.messages,
+ default_template=self.chat_template,
+ default_template_content_format=self.chat_template_content_format,
+ default_template_kwargs=self._effective_chat_template_kwargs(
+ request
+ ),
+ tool_dicts=tools,
+ parser=self.parser,
+ )
+ return _tokens(inputs[0]) if len(inputs) == 1 else None
+
+ if await render(view) != _tokens(inputs[0]):
+ return
+ if 0 in turn.outputs:
+ await _observe(
+ turn,
+ await chat_response_prefixes(
+ tokenizer,
+ view,
+ turn.prompt,
+ [(messages[0], *turn.outputs[0])],
+ render,
+ ),
+ )
+
+ token = _CURRENT.set(turn)
+ try:
+ result = await original(self, request, raw_request)
+ if not isinstance(result, AsyncIterator):
+ choices = getattr(result, "choices", ())
+ await observe(
+ {
+ c.index: c.message.model_dump(exclude_none=True)
+ for c in choices
+ if hasattr(c, "message")
+ }
+ )
+ await observe_response(result)
+ return result
+ finally:
+ _CURRENT.reset(token)
+
+ async def stream():
+ token = _CURRENT.set(turn)
+ messages = {}
+ try:
+ async for event in result:
+ if (
+ observe_chat
+ and isinstance(event, str)
+ and event.startswith("data: ")
+ ):
+ data = event[6:].strip()
+ if data == "[DONE]":
+ await observe(messages)
+ else:
+ for choice in json.loads(data).get("choices", []):
+ message = messages.setdefault(
+ choice["index"], {"role": "assistant"}
+ )
+ merge_chat_delta(message, choice.get("delta") or {})
+ if observe_responses and getattr(event, "type", None) in (
+ "response.completed",
+ "response.incomplete",
+ ):
+ await observe_response(event.response)
+ yield event
+ finally:
+ _CURRENT.reset(token)
+
+ return stream()
+
+ setattr(owner, method, create)
+
+ install(chat, "create_chat_completion", observe_chat=True)
+ install(completion, "create_completion", observe_chat=False)
+ install(
+ responses_serving.OpenAIServingResponses,
+ "create_responses",
+ observe_responses=True,
+ )
diff --git a/tests/unit/test_append_only.py b/tests/unit/test_append_only.py
new file mode 100644
index 000000000..99bff1ad5
--- /dev/null
+++ b/tests/unit/test_append_only.py
@@ -0,0 +1,195 @@
+import asyncio
+from types import SimpleNamespace
+
+from pydantic import BaseModel
+import pytest
+
+from art.token_prefix import TokenPrefixCache, apply_prefix_edits
+from art.utils.append_only import chat_prefix_observations, chat_response_prefixes
+from art_inference.append_only import (
+ output_prefix_observations,
+ patch_deepseek_renderer,
+)
+
+
+@pytest.mark.parametrize("reasoning", [None, "", "\nthought\n"])
+def test_dsv32_earlier_turn_framing_survives_a_new_user_turn(reasoning):
+ def render(index, messages, thinking_mode):
+ last_user = max(i for i, m in enumerate(messages) if m["role"] == "user")
+ thinking = thinking_mode == "thinking"
+ if messages[index]["role"] == "user":
+ return "USER" + (
+ "" if thinking and index == last_user else ""
+ )
+ if thinking and index > last_user:
+ assert messages[index]["reasoning_content"], "Missing reasoning"
+ return messages[index]["reasoning_content"] + "END"
+ return "END"
+
+ encoding = SimpleNamespace(render_message=render, thinking_end_token="")
+ preserve = True
+ patch_deepseek_renderer(encoding, lambda: preserve, prefix_only=True)
+ messages = [
+ {"role": "user"},
+ {"role": "assistant", "reasoning_content": reasoning},
+ ]
+
+ def encode(messages):
+ return "".join(
+ encoding.render_message(i, messages, "thinking")
+ for i in range(len(messages))
+ )
+
+ completed = encode(messages)
+ assert encode([*messages, {"role": "user"}]).startswith(completed)
+ if reasoning:
+ preserve = False
+ assert not encode([*messages, {"role": "user"}]).startswith(completed)
+
+
+class Tokenizer:
+ def encode(self, text):
+ return list(text.encode())
+
+ def decode(self, tokens, *, skip_special_tokens=False):
+ return bytes(tokens).decode()
+
+
+def test_many_protocol_markers_do_not_multiply_long_prompt_storage():
+ tokenizer = SimpleNamespace(
+ all_special_ids=[1, 2, 3],
+ encode=lambda text, **_: [ord(char) for char in text],
+ decode=lambda tokens, **_: "".join(chr(token) for token in tokens),
+ )
+ prompt = [100] * 128_000
+ output = [101, 1, 102, 2, 103, 3] * 100
+ entries = output_prefix_observations(tokenizer, prompt, prompt, output)
+ assert len(entries) <= 2
+ assert entries[-1][1] == prompt + output
+
+
+@pytest.mark.parametrize("action", ["pass", '{"x": 1}'])
+def test_normalized_reasoning_survives_an_edited_action(action):
+ tokenizer = Tokenizer()
+ encode = tokenizer.encode
+ prompt = encode("USER question ASSISTANT \n")
+ rendered = prompt + encode(f"reasoning\n\n\n{action}END\n")
+ sampled = encode(f"\nreasoning\n\n\n{action}END")
+ reasoning = prompt + encode("reasoning\n\n\nEND\n")
+ observations = chat_prefix_observations(
+ tokenizer, prompt, rendered, prompt, sampled, reasoning_prompt=reasoning
+ )
+ assert len(observations) == 2
+ cache = TokenPrefixCache()
+ for rendered_prefix, raw_prefix, edits in observations:
+ assert apply_prefix_edits(rendered_prefix, edits) == raw_prefix
+ cache.insert(rendered_prefix, raw_prefix, "rollout", edits)
+
+ for next_action in (action, "corrected action"):
+ next_prompt = prompt + encode(
+ f"reasoning\n\n\n{next_action}END\nUSER next ASSISTANT"
+ )
+ match = cache.lookup(next_prompt, "rollout")
+ assert match is not None
+ served = list(match.raw_prefix) + next_prompt[match.rendered_length :]
+ assert served == prompt + encode(
+ f"\nreasoning\n\n\n{next_action}END\nUSER next ASSISTANT"
+ )
+
+
+def test_truncation_does_not_replace_completed_turn_framing():
+ tokenizer = Tokenizer()
+ encode = tokenizer.encode
+ assert not chat_prefix_observations(
+ tokenizer,
+ encode("prompt"),
+ encode("prompt partial END"),
+ encode("prompt"),
+ encode(" partial"),
+ complete=False,
+ )
+
+
+def test_changed_prompt_cannot_be_registered_as_a_continuation():
+ assert not chat_prefix_observations(Tokenizer(), [1], [2, 3], [4], [5])
+
+
+@pytest.mark.parametrize(
+ "options",
+ [{"preserve_thinking": False}, {"clear_thinking": True}, {"drop_thinking": True}],
+)
+def test_explicit_rewriting_bypasses_observation(options):
+ class Request(BaseModel):
+ messages: list[dict]
+ chat_template_kwargs: dict
+
+ async def render(_):
+ pytest.fail("Explicit rewriting must not create preservation mappings")
+
+ assert (
+ asyncio.run(
+ chat_response_prefixes(
+ Tokenizer(),
+ Request(messages=[], chat_template_kwargs=options),
+ [],
+ [],
+ render,
+ )
+ )
+ == []
+ )
+
+
+def test_custom_stop_does_not_delete_the_template_terminator():
+ encode = Tokenizer().encode
+ assert (
+ chat_prefix_observations(
+ Tokenizer(),
+ encode("prompt:"),
+ encode("prompt:answerEND"),
+ encode("prompt:"),
+ encode("answer"),
+ complete=True,
+ )
+ == []
+ )
+
+
+def test_reasoning_boundary_does_not_absorb_whitespace_merged_into_action():
+ from art_inference.append_only import _whitespace_prefix
+
+ tokenizer = Tokenizer()
+ raw = tokenizer.encode("thought\n\naction")
+ boundary = _whitespace_prefix(tokenizer, tokenizer.encode("thought"), raw)
+ assert tokenizer.decode(raw[:boundary]) == "thought"
+
+
+def test_missing_lineage_never_selects_an_ambiguous_tokenization():
+ cache = TokenPrefixCache()
+ cache.insert([1, 2], [1, 3], "first")
+ match = cache.lookup([1, 2], None)
+ assert match is not None and match.raw_prefix == (1, 3)
+ cache.insert([1, 2], [1, 4], "second")
+ assert cache.lookup([1, 2], None) is None
+ first = cache.lookup([1, 2], "first")
+ second = cache.lookup([1, 2], "second")
+ assert first is not None and first.raw_prefix == (1, 3)
+ assert second is not None and second.raw_prefix == (1, 4)
+
+
+def test_history_edits_leave_multimodal_markers_outside_replacements():
+ from types import SimpleNamespace
+
+ from art_inference.append_only import _rendering_edits
+ from art_inference.vllm import replace_prompt_tokens
+
+ rendered = [1, 2, 9, 9, 3, 4]
+ raw = [1, 1, 2, 9, 9, 3, 4, 4]
+ edits = _rendering_edits(SimpleNamespace(all_special_ids=[9]), rendered, raw)
+ assert apply_prefix_edits(rendered, edits) == raw
+ original = {
+ "prompt_token_ids": rendered,
+ "mm_placeholders": {"image": [SimpleNamespace(offset=2, length=2)]},
+ }
+ updated = replace_prompt_tokens(original, raw, edits)
+ assert updated["mm_placeholders"]["image"][0].offset == 3
diff --git a/tests/unit/test_inference_history.py b/tests/unit/test_inference_history.py
new file mode 100644
index 000000000..a823a9bce
--- /dev/null
+++ b/tests/unit/test_inference_history.py
@@ -0,0 +1,401 @@
+import asyncio
+import json
+from types import SimpleNamespace
+
+from pydantic import BaseModel, model_validator
+import pytest
+
+from art_inference import vllm
+from art_inference.token_prefix import TokenPrefixStore
+
+
+class Request(BaseModel):
+ model: str = "model"
+ messages: list[dict]
+ stream: bool = False
+ stream_options: dict | None = None
+ max_tokens: int | None = None
+ max_completion_tokens: int | None = None
+ truncate_prompt_tokens: int | None = None
+ add_generation_prompt: bool = True
+ continue_final_message: bool = False
+ chat_template_kwargs: dict = {}
+ previous_response_id: str | None = None
+ tools: list = []
+ tool_choice: str = "auto"
+
+ @model_validator(mode="after")
+ def validate_stream(self):
+ if self.stream_options and not self.stream:
+ raise ValueError("Stream options can only be defined when stream=True")
+ return self
+
+
+class Message(BaseModel):
+ role: str = "assistant"
+ reasoning_content: str
+ content: str
+
+
+@pytest.fixture
+def serving(monkeypatch):
+ monkeypatch.setattr(vllm, "_PREFIXES", TokenPrefixStore())
+
+ class Tokenizer:
+ all_special_ids = [ord("#"), ord("~")]
+
+ def encode(self, text, **kwargs):
+ return list(text.encode())
+
+ def decode(self, tokens, **kwargs):
+ return bytes(tokens).decode()
+
+ class Renderer:
+ def __init__(self, config, tokenizer):
+ self.tokenizer = tokenizer
+
+ class Params:
+ def get_apply_chat_template_kwargs(self):
+ return {}
+
+ class Engine:
+ def __init__(self):
+ self.prompts = []
+
+ async def generate(self, prompt, sampling_params=None):
+ self.prompts.append(prompt["prompt_token_ids"])
+ chunks = (
+ [b"\nthought\n#", b"action~"]
+ if sampling_params
+ else [b"\nthought\n#action~"]
+ )
+ for index, tokens in enumerate(chunks):
+ yield SimpleNamespace(
+ outputs=[
+ SimpleNamespace(
+ index=0,
+ token_ids=list(tokens),
+ finish_reason="stop" if index == len(chunks) - 1 else None,
+ )
+ ]
+ )
+
+ class Serving:
+ def __init__(self):
+ self.renderer = Renderer(None, Tokenizer())
+ self.engine = Engine()
+ self.online_renderer = SimpleNamespace(preprocess_chat=self.preprocess_chat)
+ self.chat_template = None
+ self.chat_template_content_format = "string"
+ self.parser = None
+
+ async def preprocess_chat(self, request, messages, **kwargs):
+ return await self.render_chat_request(request)
+
+ def _effective_chat_template_kwargs(self, request):
+ return request.chat_template_kwargs
+
+ async def _make_request(self, request, previous):
+ return request.messages, (await self.render_chat_request(request))[1]
+
+ async def render_chat_request(self, request):
+ text = ""
+ for message in request.messages:
+ if message["role"] == "assistant":
+ text += (
+ "A"
+ + message.get("reasoning_content", "").strip()
+ + "#"
+ + (message.get("content") or "")
+ + "~\n"
+ )
+ else:
+ text += "U" + message["content"] + ";"
+ if request.add_generation_prompt:
+ text += "A"
+ return [], [{"prompt_token_ids": list(text.encode())}]
+
+ async def create_chat_completion(self, request, raw_request=None):
+ _, inputs = await self.render_chat_request(request)
+ message = Message(reasoning_content="\nthought\n", content="action")
+
+ async def stream():
+ params = SimpleNamespace(output_kind=SimpleNamespace(name="DELTA"))
+ async for _ in self.engine.generate(inputs[0], params):
+ pass
+ yield (
+ "data: "
+ + json.dumps(
+ {"choices": [{"index": 0, "delta": message.model_dump()}]}
+ )
+ + "\n\n"
+ )
+ yield "data: [DONE]\n\n"
+
+ if request.stream:
+ return stream()
+ async for _ in self.engine.generate(inputs[0]):
+ pass
+ return SimpleNamespace(choices=[SimpleNamespace(index=0, message=message)])
+
+ async def create_completion(self, request, raw_request=None):
+ return await self.create_chat_completion(request, raw_request)
+
+ async def create_responses(self, request, raw_request=None):
+ _, inputs = await self.render_chat_request(request)
+ async for _ in self.engine.generate(inputs[0]):
+ pass
+ response = SimpleNamespace(
+ output=[
+ Message(
+ reasoning_content="\nthought\n", content="action"
+ ).model_dump()
+ ]
+ )
+ if not request.stream:
+ return response
+
+ async def events():
+ yield SimpleNamespace(type="response.completed", response=response)
+
+ return events()
+
+ def construct_input_messages(**kwargs):
+ return kwargs
+
+ utils = SimpleNamespace(
+ construct_input_messages=construct_input_messages,
+ construct_chat_messages_with_tool_call=lambda items: items,
+ construct_tool_dicts=lambda *_: None,
+ )
+ modules = {
+ "vllm.renderers.base": SimpleNamespace(BaseRenderer=Renderer),
+ "vllm.renderers.params": SimpleNamespace(ChatParams=Params),
+ "vllm.v1.engine.async_llm": SimpleNamespace(AsyncLLM=Engine),
+ "vllm.entrypoints.openai.chat_completion.serving": SimpleNamespace(
+ OpenAIServingChat=Serving
+ ),
+ "vllm.entrypoints.openai.completion.serving": SimpleNamespace(
+ OpenAIServingCompletion=Serving
+ ),
+ "vllm.entrypoints.openai.responses.serving": SimpleNamespace(
+ OpenAIServingResponses=Serving
+ ),
+ "vllm.entrypoints.openai.responses.utils": utils,
+ "vllm.entrypoints.openai.parser.harmony_utils": SimpleNamespace(
+ render_for_completion=lambda messages: messages
+ ),
+ "vllm.entrypoints.openai.chat_completion.protocol": SimpleNamespace(
+ ChatCompletionRequest=Request
+ ),
+ }
+ vllm.patch_history(modules.__getitem__)
+ return Serving(), modules
+
+
+@pytest.mark.parametrize("stream", [False, True])
+@pytest.mark.parametrize("action", ["action", "edited action"])
+def test_served_history_survives_renderer_normalization(serving, stream, action):
+ server, _ = serving
+
+ async def run():
+ user = {"role": "user", "content": "question"}
+ first = Request(
+ messages=[user],
+ stream=stream,
+ stream_options={"include_usage": True} if stream else None,
+ )
+ response = await server.create_chat_completion(first)
+ if stream:
+ assert (await anext(response)).startswith("data: ")
+ assert await anext(response) == "data: [DONE]\n\n"
+ await response.aclose()
+ prior = server.engine.prompts[0] + list(b"\nthought\n#")
+ assistant = Message(
+ reasoning_content="\nthought\n", content=action
+ ).model_dump()
+ await server.create_chat_completion(
+ Request(messages=[user, assistant, {"role": "user", "content": "next"}])
+ )
+ assert server.engine.prompts[-1][: len(prior)] == prior
+ assert ("#" + action + "~\nU").encode() in bytes(server.engine.prompts[-1])
+
+ asyncio.run(run())
+
+
+@pytest.mark.parametrize("stream", [False, True])
+@pytest.mark.parametrize("budget_field", ["max_tokens", "max_completion_tokens"])
+@pytest.mark.parametrize("truncate", [None, 11])
+def test_completed_observation_has_no_second_generation_budget_or_truncation(
+ serving, stream, budget_field, truncate
+):
+ server, _ = serving
+ render = server.render_chat_request
+ completed = []
+
+ async def validated_render(request):
+ result = await render(request)
+ prompt = result[1][0]
+ if request.truncate_prompt_tokens is not None:
+ prompt["prompt_token_ids"] = prompt["prompt_token_ids"][
+ -request.truncate_prompt_tokens :
+ ]
+ output_budget = request.max_completion_tokens or request.max_tokens or 0
+ assert len(prompt["prompt_token_ids"]) + output_budget <= 48
+ if not request.add_generation_prompt:
+ completed.append(bytes(prompt["prompt_token_ids"]))
+ return result
+
+ server.render_chat_request = validated_render
+ request = Request(
+ messages=[{"role": "user", "content": "question"}],
+ stream=stream,
+ truncate_prompt_tokens=truncate,
+ max_tokens=32 if budget_field == "max_tokens" else None,
+ max_completion_tokens=32 if budget_field == "max_completion_tokens" else None,
+ )
+
+ async def run():
+ response = await server.create_chat_completion(request)
+ if stream:
+ assert [event async for event in response][-1] == "data: [DONE]\n\n"
+ else:
+ assert response.choices[0].message.content == "action"
+
+ asyncio.run(run())
+ assert b"Uquestion;Athought#action~\n" in completed
+ assert getattr(request, budget_field) == 32
+ assert request.truncate_prompt_tokens == truncate
+
+
+def test_external_observations_use_the_host_store_without_local_publication(serving):
+ server, modules = serving
+ observations = []
+
+ async def observe(request, entries):
+ observations.extend(entries)
+
+ vllm.set_external_history_observer(observe, modules.__getitem__)
+
+ async def run():
+ await server.create_chat_completion(
+ Request(messages=[{"role": "user", "content": "question"}]),
+ SimpleNamespace(headers={"x-caladan-prefix-scope": "scope"}),
+ )
+
+ asyncio.run(run())
+ assert observations
+ assert not vllm._PREFIXES._scopes
+
+
+def test_previous_responses_keep_reasoning_and_tool_calls(serving):
+ _, modules = serving
+ previous = [{"type": "reasoning"}, {"type": "function_call"}]
+ utils = modules["vllm.entrypoints.openai.responses.utils"]
+ assert utils.construct_input_messages(
+ prev_response_output=previous, request_input="next"
+ )["request_input"] == [*previous, {"role": "user", "content": "next"}]
+
+
+def test_batched_generations_keep_separate_prompt_histories(serving):
+ server, _ = serving
+ turn = vllm._Turn("scope", server.renderer.tokenizer)
+
+ async def run():
+ token = vllm._CURRENT.set(turn)
+ try:
+ first = server.engine.generate({"prompt_token_ids": list(b"first")})
+ second = server.engine.generate({"prompt_token_ids": list(b"second")})
+ await anext(first)
+ await anext(second)
+ assert [item async for item in first] == []
+ assert [item async for item in second] == []
+ finally:
+ vllm._CURRENT.reset(token)
+
+ asyncio.run(run())
+ for prompt in (b"first", b"second"):
+ tokens = list(prompt + b"\nthought\n#action~")
+ match = vllm._PREFIXES.lookup("scope", tokens, None)
+ assert match is not None
+ assert list(match.raw_prefix) == tokens
+
+
+def test_vllm_multimodal_offsets_follow_preserved_text():
+ from art_inference.token_prefix import PrefixEdit
+
+ span = SimpleNamespace(offset=4, length=2)
+ original = {
+ "prompt_token_ids": [1, 2, 3, 4, 9, 9],
+ "mm_placeholders": {"image": [span]},
+ "token_type_ids": [0, 0, 0, 0, 1, 1],
+ }
+ edited = vllm.replace_prompt_tokens(
+ original, [1, 2, 2, 3, 4, 9, 9], [PrefixEdit(1, 2, (2, 2))]
+ )
+ assert edited["mm_placeholders"]["image"][0].offset == 5
+ assert span.offset == 4
+ assert edited["token_type_ids"] == [0, 0, 0, 0, 0, 1, 1]
+ with pytest.raises(ValueError, match="multimodal placeholder"):
+ vllm.replace_prompt_tokens(original, [1], [PrefixEdit(4, 6, ())])
+
+
+def test_native_deepseek_v32_respects_preservation_options():
+ encoding = SimpleNamespace(
+ encode_messages=lambda **kwargs: kwargs,
+ render_message=lambda index, messages, **kwargs: messages[index],
+ )
+
+ class Tokenizer:
+ def apply_chat_template(self, messages, **kwargs):
+ return encoding.encode_messages(
+ drop_thinking=messages[-1]["role"] == "user"
+ )
+
+ Tokenizer.apply_chat_template.__module__ = "vllm.tokenizers.deepseek_v32"
+ tokenizer = Tokenizer()
+ vllm._configure_native_tokenizer(tokenizer, lambda _: encoding)
+ messages = [{"role": "user"}]
+ assert tokenizer.apply_chat_template(messages)["drop_thinking"] is False
+ assert (
+ tokenizer.apply_chat_template(messages, drop_thinking=True)["drop_thinking"]
+ is True
+ )
+ assert (
+ tokenizer.apply_chat_template(messages, preserve_thinking=False)[
+ "drop_thinking"
+ ]
+ is True
+ )
+ assert tokenizer.apply_chat_template(messages)["drop_thinking"] is False
+
+
+@pytest.mark.parametrize("legacy", [False, True])
+@pytest.mark.parametrize("stream", [False, True])
+def test_responses_use_the_native_renderer_and_preserve_ids(serving, legacy, stream):
+ server, _ = serving
+ if legacy:
+ server.openai_serving_render = server.online_renderer
+ del server.online_renderer
+
+ async def run():
+ request = Request(
+ messages=[{"role": "user", "content": "question"}], stream=stream
+ )
+ response = await server.create_responses(request)
+ if stream:
+ assert (await anext(response)).type == "response.completed"
+ await response.aclose()
+ previous = Request(
+ messages=[
+ *request.messages,
+ Message(reasoning_content="\nthought\n", content="action").model_dump(),
+ {"role": "user", "content": "next"},
+ ]
+ )
+ await server.create_responses(previous)
+ assert bytes(server.engine.prompts[-1]).startswith(
+ b"Uquestion;A\nthought\n#action~"
+ )
+
+ asyncio.run(run())
diff --git a/tests/unit/test_local_sft.py b/tests/unit/test_local_sft.py
index f9c423b2f..7b27eb38d 100644
--- a/tests/unit/test_local_sft.py
+++ b/tests/unit/test_local_sft.py
@@ -38,7 +38,7 @@ def test_qwen_rollout_server_uses_preserve_thinking_template(
lambda *_args: None,
)
monkeypatch.setattr(
- "art.local.backend.AutoTokenizer.from_pretrained",
+ "art.local.backend.get_tokenizer",
lambda _model: tokenizer,
)
config: dict[str, Any] = {}
@@ -49,24 +49,21 @@ def test_qwen_rollout_server_uses_preserve_thinking_template(
assert "preserve_thinking is defined and preserve_thinking is true" in configured
-def test_non_qwen_rollout_server_does_not_load_template(
- monkeypatch: pytest.MonkeyPatch,
-) -> None:
+def test_all_model_families_use_the_default_tokenizer(monkeypatch):
monkeypatch.setattr(
- "art.local.backend._model_support_default_chat_template",
- lambda *_args: None,
+ "art.local.backend._model_support_default_chat_template", lambda *_: None
)
- monkeypatch.setattr(
- "art.local.backend.AutoTokenizer.from_pretrained",
- lambda _model: pytest.fail("non-Qwen templates should not be loaded eagerly"),
- )
- config: dict[str, Any] = {}
+ calls = []
- _apply_configured_chat_template_server_args(
- config, {}, base_model="meta-llama/Llama-3.1-8B-Instruct"
- )
+ def load(model):
+ calls.append(model)
+ return type("Tokenizer", (), {"chat_template": "{{ messages }}"})()
- assert config == {}
+ monkeypatch.setattr("art.local.backend.get_tokenizer", load)
+ config = {}
+ _apply_configured_chat_template_server_args(config, {}, base_model="other/family")
+ assert calls == ["other/family"]
+ assert config["server_args"]["chat_template"] == "{{ messages }}"
def test_explicit_server_template_avoids_tokenizer_loading(
@@ -77,7 +74,7 @@ def test_explicit_server_template_avoids_tokenizer_loading(
lambda *_args: pytest.fail("explicit template should win before model support"),
)
monkeypatch.setattr(
- "art.local.backend.AutoTokenizer.from_pretrained",
+ "art.local.backend.get_tokenizer",
lambda _model: pytest.fail("explicit template should avoid hub access"),
)
config: dict[str, Any] = {"server_args": {"chat_template": "explicit"}}
@@ -97,7 +94,7 @@ def test_qwen_template_load_failure_is_a_warned_fallback(
lambda *_args: None,
)
monkeypatch.setattr(
- "art.local.backend.AutoTokenizer.from_pretrained",
+ "art.local.backend.get_tokenizer",
lambda _model: (_ for _ in ()).throw(ValueError("bad tokenizer")),
)
config: dict[str, Any] = {}
@@ -119,7 +116,7 @@ def _local_sft_patches(
with ExitStack() as stack:
for patcher in (
patch(
- "art.local.backend.AutoTokenizer.from_pretrained",
+ "art.local.backend.get_tokenizer",
return_value=object(),
),
patch.object(
@@ -257,3 +254,22 @@ async def train_sft(
assert [call["learning_rate"] for call in calls] == [1e-4, 1e-4]
assert [batch.learning_rate for batch in captured_batches] == [1e-4]
assert results[0]["data/step_num_dropped_trajectories"] == 1.0
+
+
+def test_python_encoder_marker_is_not_passed_as_a_jinja_template(monkeypatch):
+ monkeypatch.setattr(
+ "art.local.backend._model_support_default_chat_template", lambda *_: None
+ )
+ monkeypatch.setattr(
+ "art.local.backend.get_tokenizer",
+ lambda _: type(
+ "Tokenizer",
+ (),
+ {"chat_template": "deepseek_v4_python_encoder enable_thinking"},
+ )(),
+ )
+ config = {}
+ _apply_configured_chat_template_server_args(
+ config, {}, base_model="deepseek-ai/DeepSeek-V4-Flash"
+ )
+ assert "chat_template" not in config.get("server_args", {})
diff --git a/tests/unit/test_pipeline_trainer_local_backend.py b/tests/unit/test_pipeline_trainer_local_backend.py
index c59141eca..35b1753ec 100644
--- a/tests/unit/test_pipeline_trainer_local_backend.py
+++ b/tests/unit/test_pipeline_trainer_local_backend.py
@@ -746,7 +746,7 @@ def test_local_backend_get_packed_tensors_warns_and_drops_overlong_results(
with (
patch(
- "art.local.backend.AutoTokenizer.from_pretrained",
+ "art.local.backend.get_tokenizer",
return_value=short_result._tokenizer,
),
patch("transformers.AutoImageProcessor.from_pretrained", return_value=None),
diff --git a/tests/unit/test_prefix_cache.py b/tests/unit/test_prefix_cache.py
deleted file mode 100644
index 9c4b0ba65..000000000
--- a/tests/unit/test_prefix_cache.py
+++ /dev/null
@@ -1,36 +0,0 @@
-"""Tests for the LRUTrieCache prefix rewrite helper."""
-
-import pytest
-
-pytest.importorskip("datrie")
-
-from art.tinker.prefix_cache import LRUTrieCache
-
-
-class TestLRUTrieCache:
- def test_longest_prefix_match(self) -> None:
- cache = LRUTrieCache(max_entries=10)
- cache.insert([1, 2], [10, 11])
- cache.insert([1, 2, 3], [20, 21, 22])
-
- entry = cache.lookup([1, 2, 3, 4])
-
- assert entry is not None
- assert entry.rendered_len == 3
- assert entry.raw_prefix == (20, 21, 22)
-
- def test_lru_eviction(self) -> None:
- cache = LRUTrieCache(max_entries=2)
- cache.insert([1], [10])
- cache.insert([2], [20])
-
- assert cache.lookup([1, 99]) is not None
-
- cache.insert([3], [30])
-
- assert cache.lookup([2, 0]) is None
- assert cache.lookup([1, 0]) is not None
-
- def test_invalid_size(self) -> None:
- with pytest.raises(ValueError):
- LRUTrieCache(max_entries=0)
diff --git a/tests/unit/test_preprocessing_tokenize.py b/tests/unit/test_preprocessing_tokenize.py
index 40bc11349..1720bdffc 100644
--- a/tests/unit/test_preprocessing_tokenize.py
+++ b/tests/unit/test_preprocessing_tokenize.py
@@ -441,7 +441,8 @@ def test_native_or_unrecognized_thinking_templates_are_unchanged() -> None:
)
unrelated = "{%- if loop.index0 > ns.last_query_index %}content{% endif %}"
- assert chat_template_with_preserved_thinking(native) == native
+ configured = chat_template_with_preserved_thinking(native)
+ assert chat_template_with_preserved_thinking(configured) == configured
assert chat_template_with_preserved_thinking(unrelated) == unrelated
diff --git a/tests/unit/test_sglang_history.py b/tests/unit/test_sglang_history.py
new file mode 100644
index 000000000..391a2d80d
--- /dev/null
+++ b/tests/unit/test_sglang_history.py
@@ -0,0 +1,339 @@
+import asyncio
+import json
+from types import SimpleNamespace
+
+from pydantic import BaseModel
+import pytest
+
+from art_inference import sglang
+from art_inference.token_prefix import TokenPrefixStore
+
+
+class Request(BaseModel):
+ model: str = "model"
+ messages: list[dict] = []
+ input: str | list = "next"
+ chat_template_kwargs: dict = {}
+ add_generation_prompt: bool = True
+ skip_special_tokens: bool = True
+ previous_response_id: str | None = None
+ stream: bool = False
+ continue_final_message: bool = False
+
+
+class Response(BaseModel):
+ id: str = "previous"
+ output: list = []
+
+
+@pytest.fixture
+def serving():
+ class Tokenizer:
+ def encode(self, text, **kwargs):
+ return list(text.encode())
+
+ def decode(self, tokens, **kwargs):
+ return bytes(tokens).decode()
+
+ def apply_chat_template(
+ self, messages, *, add_generation_prompt=True, **kwargs
+ ):
+ text = "".join(
+ "A"
+ + m.get("reasoning_content", "").strip()
+ + "#"
+ + m.get("content", "")
+ + "~\n"
+ if m["role"] == "assistant"
+ else "U" + m["content"] + ";"
+ for m in messages
+ )
+ return self.encode(text + ("A" if add_generation_prompt else ""))
+
+ encoding = SimpleNamespace(
+ encode_messages=lambda **kwargs: kwargs,
+ render_message=lambda index, messages, thinking_mode: messages[index],
+ )
+
+ class Serving:
+ def __init__(self):
+ self.tokenizer_manager = SimpleNamespace(
+ tokenizer=Tokenizer(), model_config=SimpleNamespace(is_multimodal=False)
+ )
+ self.msg_store = {}
+ self.response_store = {}
+
+ def _process_messages(self, request, is_multimodal):
+ request.skip_special_tokens = False
+ self.options = request.chat_template_kwargs
+ self.dsv4_options = encoding.encode_messages()
+ messages, _ = self._handle_last_assistant_message(
+ [dict(message) for message in request.messages], request
+ )
+ return SimpleNamespace(
+ prompt_ids=self.tokenizer_manager.tokenizer.apply_chat_template(
+ messages
+ )
+ )
+
+ def _handle_last_assistant_message(self, messages, request):
+ # SGLang 0.5.15 request rendering rewrites a trailing assistant.
+ if messages and messages[-1]["role"] == "assistant":
+ if request.continue_final_message:
+ return messages[:-1], messages[-1]["content"]
+ messages[-1] = {"role": "user", "content": messages[-1]["content"]}
+ return messages, None
+
+ async def _handle_non_streaming_request(
+ self, adapted_request, request, raw_request
+ ):
+ prompt = self._process_messages(request, False).prompt_ids
+ choice = {
+ "index": 0,
+ "message": {
+ "role": "assistant",
+ "reasoning_content": "\nthought\n",
+ "content": "action",
+ },
+ "token_ids": list(b"\nthought\n#action~"),
+ "finish_reason": "stop",
+ }
+ return SimpleNamespace(
+ body=json.dumps(
+ {"prompt_token_ids": prompt, "choices": [choice]}
+ ).encode()
+ )
+
+ async def _generate_chat_stream(self, adapted_request, request, raw_request):
+ result = json.loads(
+ (
+ await self._handle_non_streaming_request(
+ adapted_request, request, raw_request
+ )
+ ).body
+ )
+ choice = result["choices"][0]
+ choice["delta"] = choice.pop("message")
+ yield "data: " + json.dumps(result) + "\n\n"
+ yield "data: [DONE]\n\n"
+
+ async def create_responses(self, request, raw_request=None):
+ prompt = self._process_messages(
+ Request(messages=request.input), False
+ ).prompt_ids
+ payload = {
+ "status": "completed",
+ "output": [
+ {"type": "reasoning", "content": [{"text": "\nthought\n"}]},
+ {"role": "assistant", "content": "action"},
+ ],
+ "token_generations": [
+ {
+ "prompt_token_ids": prompt,
+ "output_tokens": [
+ {"token_id": token} for token in b"\nthought\n#action~"
+ ],
+ }
+ ],
+ }
+ if not request.stream:
+ return SimpleNamespace(body=json.dumps(payload).encode())
+
+ async def events():
+ yield (
+ "event: response.completed\ndata: "
+ + json.dumps({"response": payload})
+ + "\n\n"
+ )
+
+ return events()
+
+ @staticmethod
+ def _merge_consecutive_assistant_messages(messages):
+ combined = {}
+ for message in messages:
+ combined.update(message)
+ return [combined]
+
+ def _response_tools_to_chat_tools(self, request):
+ return []
+
+ @classmethod
+ def _normalize_response_message_for_chat(cls, message):
+ return message
+
+ def _construct_input_messages(self, request, prev_response=None):
+ return (
+ [*prev_response.output, *request.input]
+ if prev_response
+ else request.input
+ )
+
+ def _construct_input_messages_with_harmony(self, request, prev_response):
+ previous = self.msg_store[prev_response.id]
+ previous[:] = [m for m in previous if m.channel != "analysis"]
+ return [*previous, request.input]
+
+ entries = []
+
+ async def observe(request, observations):
+ entries.extend(observations)
+
+ modules = {
+ "sglang.srt.entrypoints.openai.serving_chat": SimpleNamespace(
+ OpenAIServingChat=Serving
+ ),
+ "sglang.srt.entrypoints.openai.serving_responses": SimpleNamespace(
+ OpenAIServingResponses=Serving
+ ),
+ "sglang.srt.entrypoints.openai.encoding_dsv4": encoding,
+ "sglang.srt.entrypoints.openai.protocol": SimpleNamespace(
+ ChatCompletionRequest=Request
+ ),
+ "sglang.srt.entrypoints.harmony_utils": SimpleNamespace(
+ render_for_completion=lambda messages: messages
+ ),
+ "sglang.srt.entrypoints.openai.encoding_dsv32": SimpleNamespace(
+ encode_messages=lambda **kwargs: kwargs,
+ render_message=lambda index, messages, **kwargs: messages[index],
+ ),
+ }
+ sglang.patch_history(observe, modules.__getitem__)
+ return Serving(), entries
+
+
+@pytest.mark.parametrize("stream", [False, True])
+def test_sglang_observes_normalized_history_and_edited_actions(serving, stream):
+ server, entries = serving
+ request = Request(messages=[{"role": "user", "content": "question"}])
+
+ async def run():
+ if stream:
+ async for _ in server._generate_chat_stream(None, request, None):
+ pass
+ else:
+ await server._handle_non_streaming_request(None, request, None)
+
+ asyncio.run(run())
+ store = TokenPrefixStore()
+ for rendered, raw, edits in entries:
+ store.insert("scope", rendered, raw, "lineage", edits)
+ prompt = list(b"Uquestion;Athought#edited~\nUnext;A")
+ match = store.lookup("scope", prompt, "lineage")
+ assert match is not None
+ actual = list(match.raw_prefix) + prompt[match.rendered_length :]
+ assert bytes(actual) == b"Uquestion;A\nthought\n#edited~\nUnext;A"
+ assert server.options == {}
+ assert request.skip_special_tokens is False
+ assert server.dsv4_options["drop_thinking"] is False
+
+
+@pytest.mark.parametrize("preserve", [False, True])
+def test_harmony_opt_out_does_not_mutate_stored_history(serving, preserve):
+ server, _ = serving
+ thought, final = (
+ SimpleNamespace(channel="analysis"),
+ SimpleNamespace(channel="final"),
+ )
+ server.msg_store["previous"] = [thought, final]
+ result = server._construct_input_messages_with_harmony(
+ Request(chat_template_kwargs={"preserve_thinking": preserve}), Response()
+ )
+ assert result == ([thought, final, "next"] if preserve else [final, "next"])
+ assert server.msg_store["previous"] == [thought, final]
+
+
+def test_responses_use_full_reasoning_and_retain_output_items(serving):
+ server, _ = serving
+ assert server._normalize_response_message_for_chat(
+ {
+ "type": "reasoning",
+ "content": [{"text": "full "}, {"text": "reasoning"}],
+ "summary": [{"text": "short summary"}],
+ }
+ ) == {"role": "assistant", "reasoning_content": "full reasoning"}
+ items = [{"type": "reasoning"}, {"type": "function_call"}]
+ assert server._construct_input_messages(Request(), Response(output=items)) == [
+ *items,
+ {"role": "user", "content": "next"},
+ ]
+
+
+def test_explicit_dsv4_opt_out_is_scoped_to_one_request(serving):
+ server, _ = serving
+ server._process_messages(
+ Request(chat_template_kwargs={"drop_thinking": True}), False
+ )
+ assert server.dsv4_options["drop_thinking"] is True
+ server._process_messages(Request(), False)
+ assert server.dsv4_options["drop_thinking"] is False
+
+
+@pytest.mark.parametrize("stream", [False, True])
+def test_responses_history_opt_out_reaches_nested_chat_rendering(serving, stream):
+ server, entries = serving
+
+ async def run():
+ result = await server.create_responses(
+ Request(
+ input=[{"role": "user", "content": "question"}],
+ chat_template_kwargs={"preserve_thinking": False},
+ stream=stream,
+ )
+ )
+ if stream:
+ async for _ in result:
+ pass
+
+ asyncio.run(run())
+ assert server.dsv4_options["drop_thinking"] is True
+ assert not entries
+ server._process_messages(Request(), False)
+ assert server.dsv4_options["drop_thinking"] is False
+
+
+def test_multimodal_offsets_and_rope_positions_follow_text_edits():
+ import torch
+
+ from art_inference.token_prefix import PrefixEdit
+
+ mm = SimpleNamespace(
+ mm_items=[SimpleNamespace(offsets=[(4, 5)])],
+ input_ids=[1, 2, 3, 4, 9, 9],
+ padded_input_ids=[1, 2, 3, 4, -1, -1],
+ mrope_positions=torch.tensor(
+ [[0, 1, 2, 3, 4, 4], [0, 1, 2, 3, 4, 5], [0, 1, 2, 3, 4, 4]]
+ ),
+ mrope_position_delta=torch.tensor([0]),
+ token_type_ids=torch.tensor([0, 0, 0, 0, 1, 1]),
+ )
+ tokenized = SimpleNamespace(mm_inputs=mm)
+ sglang.replace_prompt_tokens(tokenized, [PrefixEdit(1, 2, (2, 2))])
+ actual = tokenized.mm_inputs
+ assert actual.mm_items[0].offsets == [(5, 6)]
+ assert mm.mm_items[0].offsets == [(4, 5)]
+ assert actual.padded_input_ids == [1, 2, 2, 3, 4, -1, -1]
+ assert actual.mrope_positions.tolist() == [
+ [0, 1, 2, 3, 4, 5, 5],
+ [0, 1, 2, 3, 4, 5, 6],
+ [0, 1, 2, 3, 4, 5, 5],
+ ]
+ assert actual.token_type_ids.tolist() == [0, 0, 0, 0, 0, 1, 1]
+ assert actual.mrope_position_delta is mm.mrope_position_delta
+
+
+@pytest.mark.parametrize("stream", [False, True])
+def test_responses_observe_the_complete_native_generation(serving, stream):
+ server, entries = serving
+
+ async def run():
+ result = await server.create_responses(
+ Request(input=[{"role": "user", "content": "question"}], stream=stream)
+ )
+ if stream:
+ assert (await anext(result)).startswith("event: response.completed")
+ await result.aclose()
+
+ asyncio.run(run())
+ assert entries
+ assert bytes(entries[-1][1]) == b"Uquestion;A\nthought\n#action~"
diff --git a/tests/unit/test_tinker_renderers.py b/tests/unit/test_tinker_renderers.py
index b2f641a70..81e2bed8f 100644
--- a/tests/unit/test_tinker_renderers.py
+++ b/tests/unit/test_tinker_renderers.py
@@ -1,6 +1,8 @@
+import asyncio
import json
from pathlib import Path
-from typing import cast
+from types import SimpleNamespace
+from typing import Any, cast
from tinker import EncodedTextChunk, ModelInput
from tinker_cookbook import renderers
@@ -75,6 +77,32 @@ def test_get_renderer_name_autodetects_qwen3_5() -> None:
assert get_renderer_name("Qwen/Qwen3.5-35B-A3B") == "qwen3_5_disable_thinking"
+def test_tinker_normalizes_tool_arguments_for_mapping_templates():
+ from art.tinker.server import OpenAICompatibleTinkerServerWorker
+
+ class Tokenizer:
+ chat_template = (
+ "{% for k, v in tool_call.function.arguments.items() %}{{ k }}{% endfor %}"
+ )
+
+ def apply_chat_template(self, messages, **kwargs):
+ assert messages[0]["tool_calls"][0]["function"]["arguments"] == {"x": 1}
+ return [1, 2]
+
+ worker = OpenAICompatibleTinkerServerWorker(
+ _renderers={"model": cast(Any, SimpleNamespace(tokenizer=Tokenizer()))}
+ )
+ message = {
+ "role": "assistant",
+ "tool_calls": [{"function": {"name": "lookup", "arguments": '{"x": 1}'}}],
+ }
+ assert asyncio.run(worker.prompt_tokens("model", cast(Any, [message]), None)) == [
+ 1,
+ 2,
+ ]
+ assert message["tool_calls"][0]["function"]["arguments"] == '{"x": 1}'
+
+
def test_qwen3_5_generation_prompt_matches_hf_suffixes() -> None:
tokenizer = FakeTokenizer()
diff --git a/tests/unit/test_token_prefix.py b/tests/unit/test_token_prefix.py
new file mode 100644
index 000000000..c4457dbcf
--- /dev/null
+++ b/tests/unit/test_token_prefix.py
@@ -0,0 +1,1074 @@
+import asyncio
+from collections import defaultdict
+from collections.abc import Awaitable, Callable
+from dataclasses import replace
+import json
+import random
+
+import pytest
+
+from art.token_prefix import (
+ COMPACT_PREFIX_VERSION,
+ CompactPrefixStore,
+ PrefixEdit,
+ PrefixMatch,
+ TokenPrefixCache,
+ TokenPrefixRuntime,
+ TokenPrefixStore,
+ apply_prefix_edits,
+ compact_candidate_from_payload,
+ compact_candidate_payload,
+ compact_candidates_from_header,
+ compact_candidates_header,
+ compact_prefix_candidate,
+ prefix_edits,
+ resolve_compact_prefix,
+ token_digest,
+ token_prefix_digests,
+)
+
+_SharedEntry = tuple[str, list[int], list[int], str, tuple[PrefixEdit, ...]]
+
+
+def _compact_entry(
+ rendered: list[int],
+ raw: list[int],
+ *,
+ scope: str = "a" * 64,
+ lineage: str = "rollout",
+) -> dict[str, object]:
+ return {
+ "scope": scope,
+ "lineage": lineage,
+ **compact_candidate_payload(compact_prefix_candidate(rendered, raw)),
+ }
+
+
+def _batch(
+ insert: Callable[[str, list[int], list[int], str], Awaitable[object]],
+) -> Callable[[list[_SharedEntry]], Awaitable[None]]:
+ async def insert_many(entries: list[_SharedEntry]) -> None:
+ for scope, rendered, raw, lineage, _ in entries:
+ await insert(scope, rendered, raw, lineage)
+
+ return insert_many
+
+
+async def _discard(_entries: list[_SharedEntry]) -> None:
+ pass
+
+
+async def _runtime_insert(
+ runtime: TokenPrefixRuntime,
+ scope: str,
+ rendered: list[int],
+ raw: list[int],
+ lineage: str = "rollout",
+) -> None:
+ await runtime.insert_many(
+ [(scope, rendered, raw, lineage, prefix_edits(rendered, raw))]
+ )
+
+
+def test_store_retains_only_a_complete_bounded_lineage_snapshot() -> None:
+ store = CompactPrefixStore(max_serialized_bytes_per_lineage=512)
+ candidates = [compact_prefix_candidate([index], [index + 1]) for index in range(4)]
+ for candidate in candidates:
+ store.insert("a" * 64, "rollout", candidate)
+
+ retained = store.candidates("a" * 64, "rollout", 4)
+ payload = {
+ "version": COMPACT_PREFIX_VERSION,
+ "candidates": [compact_candidate_payload(value) for value in retained],
+ }
+ assert len(json.dumps(payload, separators=(",", ":")).encode()) <= 512
+ assert 0 < len(retained) < len(candidates)
+ assert retained[0] == candidates[-1]
+
+
+def test_default_store_retains_busy_lineage_turns() -> None:
+ store = CompactPrefixStore()
+ candidates = [
+ compact_prefix_candidate([index], [index + 1]) for index in range(240)
+ ]
+ for candidate in candidates:
+ assert store.insert("a" * 64, "shared-lineage", candidate)
+
+ retained = store.candidates("a" * 64, "shared-lineage", 240)
+ assert len(retained) == len(candidates)
+ assert candidates[0] in retained
+
+
+def test_complete_candidate_headers_round_trip_or_fall_back() -> None:
+ candidate = compact_prefix_candidate([1, 500], [1, 101, 102])
+
+ encoded = compact_candidates_header([candidate])
+ empty = compact_candidates_header([])
+ dense = compact_candidates_header(
+ [compact_prefix_candidate([index], [index]) for index in range(65)]
+ )
+ oversized = compact_candidates_header(
+ [
+ compact_prefix_candidate(
+ [index],
+ [index, *range(1_000 + index * 1_000, 2_000 + index * 1_000)],
+ )
+ for index in range(32)
+ ]
+ )
+
+ assert encoded is not None
+ assert compact_candidates_from_header(encoded) == (candidate,)
+ assert empty is not None
+ assert compact_candidates_from_header(empty) == ()
+ assert dense is not None
+ assert len(compact_candidates_from_header(dense)) == 65
+
+ assert oversized is None
+
+
+def test_candidate_header_rejects_malformed_or_partial_snapshots() -> None:
+ for value in (
+ "not-json",
+ json.dumps({"version": "old", "candidates": []}),
+ json.dumps({"version": COMPACT_PREFIX_VERSION}),
+ ):
+ try:
+ compact_candidates_from_header(value)
+ except ValueError:
+ pass
+ else:
+ raise AssertionError("invalid candidate header was accepted")
+
+
+def test_compact_edits_cover_deletions_expansions_and_sparse_128k_changes() -> None:
+ cases = (
+ ([9, 1, 2], [1, 2]),
+ ([1, 2], [9, 1, 2]),
+ ([1, 2, 3], [1, 8, 9, 3]),
+ )
+ for rendered, raw in cases:
+ edits = prefix_edits(rendered, raw)
+ assert apply_prefix_edits(rendered, edits) == raw
+ candidate = compact_prefix_candidate(rendered, raw, edits)
+ assert resolve_compact_prefix([*rendered, 77], [candidate]) == PrefixMatch(
+ len(rendered), tuple(raw)
+ )
+
+ rendered = list(range(128_000))
+ raw = rendered.copy()
+ raw[10] = 200_010
+ raw[64_000] = 264_000
+ raw[-2] = 327_998
+ edits = prefix_edits(rendered, raw)
+ payload = compact_candidate_payload(compact_prefix_candidate(rendered, raw, edits))
+
+ assert len(edits) == 3
+ assert len(json.dumps(payload, separators=(",", ":")).encode()) < 512
+ assert resolve_compact_prefix(
+ [*rendered, 9], [compact_candidate_from_payload(payload)]
+ ) == PrefixMatch(len(rendered), tuple(raw))
+
+
+def test_compact_resolution_is_longest_newest_and_corruption_tolerant() -> None:
+ shorter = compact_prefix_candidate([1, 2], [7, 8])
+ longer = compact_prefix_candidate([1, 2, 3], [7, 8, 9])
+ corrupt = replace(longer, raw_digest="0" * 64)
+
+ assert resolve_compact_prefix(
+ [1, 2, 3, 4], [shorter, corrupt, longer]
+ ) == PrefixMatch(3, (7, 8, 9))
+
+ first = compact_prefix_candidate([1, 2], [10, 11])
+ newest = compact_prefix_candidate([1, 2], [20, 21])
+ assert resolve_compact_prefix([1, 2, 3], [newest, first]) == PrefixMatch(
+ 2, (20, 21)
+ )
+
+
+def test_compact_identity_candidate_shadows_a_shorter_divergence() -> None:
+ store = CompactPrefixStore()
+ store.insert(
+ "scope",
+ "lineage",
+ compact_prefix_candidate([1, 2], [7, 8]),
+ )
+ store.insert(
+ "scope",
+ "lineage",
+ compact_prefix_candidate([1, 2, 3], [1, 2, 3]),
+ )
+
+ candidates = store.candidates("scope", "lineage", 4)
+ assert candidates[0].edits == ()
+ assert resolve_compact_prefix([1, 2, 3, 4], candidates) == PrefixMatch(3, (1, 2, 3))
+
+
+def test_compact_payload_rejects_invalid_token_and_edit_bounds() -> None:
+ payload = compact_candidate_payload(compact_prefix_candidate([1], [1]))
+ for edit in (
+ [[0, 1, [True]]],
+ [[0, 1, [-1]]],
+ [[0, 1, [0x1_0000_0000]]],
+ [[1, 1, []], [0, 1, []]],
+ ):
+ invalid = payload | {"edits": edit}
+ try:
+ compact_candidate_from_payload(invalid)
+ except ValueError:
+ pass
+ else:
+ raise AssertionError(f"accepted invalid edit: {edit}")
+
+ try:
+ compact_candidate_from_payload(payload | {"rendered_length": 262_145})
+ except ValueError:
+ pass
+ else:
+ raise AssertionError("accepted an unreachable rendered length")
+
+
+def test_compact_prefix_hashes_requested_lengths_consistently() -> None:
+ values = list(range(128_000))
+ lengths = [128_000, 32, 64_000, 32]
+ digests = token_prefix_digests(values, lengths)
+
+ assert list(digests) == [32, 64_000, 128_000]
+ for length, digest in digests.items():
+ assert (
+ digest
+ == compact_prefix_candidate(
+ values[:length], values[:length]
+ ).rendered_digest
+ )
+
+
+def test_compact_digest_encoding_has_a_stable_golden_vector() -> None:
+ assert token_digest([]) == (
+ "ff127bd383d2204e4e5509b3c4001b75988e0256b4116ffb904e0864848b04c1"
+ )
+ assert token_digest([0, 1, 0xFFFFFFFF]) == (
+ "4cc20a06621a8b2cec34aef4b495b6439931b0faa260e9308a7298cfe9e862b5"
+ )
+
+
+def test_compact_candidates_match_a_randomized_reference() -> None:
+ randomizer = random.Random(1)
+ for _ in range(1_000):
+ rendered = [
+ randomizer.randrange(32) for _ in range(randomizer.randrange(1, 65))
+ ]
+ raw = rendered.copy()
+ start = randomizer.randrange(len(raw) + 1)
+ stop = randomizer.randrange(start, len(raw) + 1)
+ raw[start:stop] = [
+ randomizer.randrange(32) for _ in range(randomizer.randrange(0, 12))
+ ]
+ suffix = [randomizer.randrange(32) for _ in range(randomizer.randrange(0, 12))]
+ candidate = compact_prefix_candidate(rendered, raw)
+
+ assert resolve_compact_prefix([*rendered, *suffix], [candidate]) == PrefixMatch(
+ len(rendered), tuple(raw)
+ )
+
+
+def test_longest_observed_prefix_rewrites_only_rendered_prefix() -> None:
+ cache = TokenPrefixCache()
+ cache.insert([1, 500], [1, 101, 102], "rollout")
+ cache.insert([1, 500, 3, 4], [1, 101, 102, 3, 4], "rollout")
+
+ assert cache.lookup([1, 500, 3, 4, 5], "rollout") == PrefixMatch(
+ rendered_length=4,
+ raw_prefix=(1, 101, 102, 3, 4),
+ )
+
+
+def test_cache_matches_a_reference_model_across_random_prefixes() -> None:
+ randomizer = random.Random(0)
+ cache = TokenPrefixCache(
+ max_variants=10_000,
+ max_token_ids=1_000_000,
+ max_lineages_per_variant=100,
+ )
+ reference: dict[tuple[int, ...], dict[tuple[int, ...], dict[str, int]]] = (
+ defaultdict(lambda: defaultdict(dict))
+ )
+
+ for observation in range(2_000):
+ rendered = tuple(
+ randomizer.randrange(8) for _ in range(randomizer.randrange(1, 12))
+ )
+ raw = tuple(randomizer.randrange(8) for _ in range(randomizer.randrange(1, 12)))
+ lineage = f"lineage-{randomizer.randrange(12)}"
+ cache.insert(rendered, raw, lineage)
+ reference[rendered][raw][lineage] = observation
+
+ query = tuple(
+ randomizer.randrange(8) for _ in range(randomizer.randrange(1, 16))
+ )
+ query_lineage = f"lineage-{randomizer.randrange(12)}"
+ expected = None
+ candidates: list[tuple[int, ...]] = sorted(
+ (
+ prefix
+ for prefix in reference
+ if len(prefix) <= len(query) and query[: len(prefix)] == prefix
+ ),
+ key=lambda prefix: len(prefix),
+ reverse=True,
+ )
+ for prefix in candidates:
+ variants = reference[prefix]
+ matching = [
+ (lineages[query_lineage], variant)
+ for variant, lineages in variants.items()
+ if query_lineage in lineages
+ ]
+ if not matching:
+ continue
+ _, selected = max(matching)
+ expected = PrefixMatch(len(prefix), selected)
+ break
+
+ assert cache.lookup(query, query_lineage) == expected
+
+
+def test_runtime_preserves_the_entire_prompt_on_a_miss() -> None:
+ async def exercise() -> None:
+ runtime = TokenPrefixRuntime(
+ lookup_shared=lambda _scope, _rendered, _lineage: asyncio.sleep(
+ 0, result=None
+ ),
+ insert_shared_many=_discard,
+ )
+ prompt = [1, 2, 3, 4]
+
+ canonical, engine_input, _ = await runtime.rewrite_with_edits(
+ "model", prompt, "rollout", shared_candidate=True
+ )
+
+ assert canonical == prompt
+ assert engine_input == prompt
+ assert canonical is not prompt
+ assert engine_input is not canonical
+
+ asyncio.run(exercise())
+
+
+def test_runtime_changes_exactly_the_matched_prefix() -> None:
+ async def exercise() -> None:
+ shared = TokenPrefixStore()
+
+ async def lookup(scope, rendered, lineage):
+ return shared.lookup(scope, rendered, lineage)
+
+ async def insert(scope, rendered, raw, lineage):
+ shared.insert(scope, rendered, raw, lineage)
+
+ runtime = TokenPrefixRuntime(
+ lookup_shared=lookup, insert_shared_many=_batch(insert)
+ )
+ cases = (
+ ([500], [101, 102], []),
+ ([1, 500], [1, 101, 102], [3]),
+ ([1, 2, 3, 4], [7], [5, 6, 7, 8]),
+ (list(range(128)), [9, 10, 11], [200, 201]),
+ )
+ for index, (rendered_prefix, raw_prefix, suffix) in enumerate(cases):
+ scope = f"model-{index}"
+ await _runtime_insert(runtime, scope, rendered_prefix, raw_prefix)
+ await runtime.flush()
+ prompt = [*rendered_prefix, *suffix]
+
+ canonical, engine_input, _ = await runtime.rewrite_with_edits(
+ scope, prompt, "rollout", shared_candidate=True
+ )
+
+ assert canonical == prompt
+ assert engine_input == [*raw_prefix, *suffix]
+ assert engine_input[len(raw_prefix) :] == suffix
+
+ asyncio.run(exercise())
+
+
+def test_replication_is_async_while_local_continuations_are_immediate() -> None:
+ async def exercise() -> None:
+ started = asyncio.Event()
+ release = asyncio.Event()
+ shared = TokenPrefixStore()
+
+ async def insert(scope, rendered, raw, lineage):
+ started.set()
+ await release.wait()
+ shared.insert(scope, rendered, raw, lineage)
+
+ runtime = TokenPrefixRuntime(
+ lookup_shared=lambda scope, rendered, lineage: asyncio.sleep(
+ 0, result=shared.lookup(scope, rendered, lineage)
+ ),
+ insert_shared_many=_batch(insert),
+ )
+ await _runtime_insert(runtime, "model", [1, 500], [1, 101, 102])
+
+ assert not started.is_set()
+ continuation = asyncio.create_task(
+ runtime.rewrite_with_edits(
+ "model",
+ [1, 500, 3],
+ "rollout",
+ shared_candidate=True,
+ )
+ )
+ _, rewritten, _ = await continuation
+ assert rewritten == [1, 101, 102, 3]
+ await started.wait()
+ release.set()
+ assert await runtime.close()
+
+ asyncio.run(exercise())
+
+
+def test_completed_replication_does_not_delay_a_later_continuation() -> None:
+ async def exercise() -> None:
+ shared = TokenPrefixStore()
+
+ async def insert(scope, rendered, raw, lineage):
+ shared.insert(scope, rendered, raw, lineage)
+
+ runtime = TokenPrefixRuntime(
+ lookup_shared=lambda scope, rendered, lineage: asyncio.sleep(
+ 0, result=shared.lookup(scope, rendered, lineage)
+ ),
+ insert_shared_many=_batch(insert),
+ )
+ await _runtime_insert(runtime, "model", [1, 500], [1, 101, 102])
+ assert await runtime.flush()
+ assert runtime._replicator is None
+
+ _, rewritten, _ = await asyncio.wait_for(
+ runtime.rewrite_with_edits(
+ "model",
+ [1, 500, 3],
+ "rollout",
+ shared_candidate=True,
+ ),
+ 0.1,
+ )
+ assert rewritten == [1, 101, 102, 3]
+ await runtime.close()
+
+ asyncio.run(exercise())
+
+
+def test_failed_replication_can_reuse_an_older_matching_shared_prefix() -> None:
+ async def exercise() -> None:
+ shared = TokenPrefixStore()
+ shared.insert("model", [1, 500], [1, 41, 42], "rollout")
+
+ async def fail(_scope, _rendered, _raw, _lineage):
+ raise RuntimeError("unavailable")
+
+ first = TokenPrefixRuntime(
+ lookup_shared=lambda scope, rendered, lineage: asyncio.sleep(
+ 0, result=shared.lookup(scope, rendered, lineage)
+ ),
+ insert_shared_many=_batch(fail),
+ )
+ await _runtime_insert(first, "model", [1, 500], [1, 101, 102])
+ assert await first.flush()
+ second = TokenPrefixRuntime(
+ lookup_shared=lambda scope, rendered, lineage: asyncio.sleep(
+ 0, result=shared.lookup(scope, rendered, lineage)
+ ),
+ insert_shared_many=_discard,
+ )
+
+ canonical, rewritten, _ = await second.rewrite_with_edits(
+ "model", [1, 500, 3], "rollout", shared_candidate=True
+ )
+ assert canonical == [1, 500, 3]
+ assert rewritten == [1, 41, 42, 3]
+ await first.close()
+ await second.close()
+
+ asyncio.run(exercise())
+
+
+def test_replication_queue_is_bounded_and_drops_only_shared_work() -> None:
+ async def exercise() -> None:
+ started = asyncio.Event()
+ release = asyncio.Event()
+ replicated: list[str] = []
+
+ async def insert(scope, _rendered, _raw, _lineage):
+ replicated.append(scope)
+ if scope == "first":
+ started.set()
+ await release.wait()
+
+ runtime = TokenPrefixRuntime(
+ lookup_shared=lambda _scope, _rendered, _lineage: asyncio.sleep(
+ 0, result=None
+ ),
+ insert_shared_many=_batch(insert),
+ max_pending_batches=1,
+ )
+ await _runtime_insert(runtime, "first", [1], [11])
+ await started.wait()
+ await _runtime_insert(runtime, "second", [2], [22])
+ await _runtime_insert(runtime, "dropped", [3], [33])
+
+ _, rewritten, _ = await runtime.rewrite_with_edits(
+ "dropped", [3, 4], "rollout", shared_candidate=False
+ )
+ assert rewritten == [33, 4]
+
+ release.set()
+ assert await runtime.close()
+ assert replicated == ["first", "second"]
+
+ asyncio.run(exercise())
+
+
+def test_shutdown_waits_for_replication_then_stops_the_worker() -> None:
+ async def exercise() -> None:
+ started = asyncio.Event()
+ release = asyncio.Event()
+
+ async def insert(_scope, _rendered, _raw, _lineage):
+ started.set()
+ await release.wait()
+
+ runtime = TokenPrefixRuntime(
+ lookup_shared=lambda _scope, _rendered, _lineage: asyncio.sleep(
+ 0, result=None
+ ),
+ insert_shared_many=_batch(insert),
+ )
+ await _runtime_insert(runtime, "model", [1], [2])
+ await started.wait()
+ closing = asyncio.create_task(runtime.close(timeout=1))
+ await asyncio.sleep(0)
+ assert not closing.done()
+
+ release.set()
+ assert await closing
+
+ asyncio.run(exercise())
+
+
+def test_replication_coalesces_and_shutdown_timeout_is_bounded() -> None:
+ async def exercise() -> None:
+ calls = []
+
+ async def insert_many(entries):
+ calls.append(entries)
+
+ runtime = TokenPrefixRuntime(
+ lookup_shared=lambda _scope, _rendered, _lineage: asyncio.sleep(
+ 0, result=None
+ ),
+ insert_shared_many=insert_many,
+ )
+ await _runtime_insert(runtime, "first", [1], [11])
+ await _runtime_insert(runtime, "second", [2], [22])
+ assert await runtime.flush()
+ assert len(calls) == 1
+ assert [entry[0] for entry in calls[0]] == ["first", "second"]
+
+ blocked = TokenPrefixRuntime(
+ lookup_shared=lambda _scope, _rendered, _lineage: asyncio.sleep(
+ 0, result=None
+ ),
+ insert_shared_many=_batch(lambda *_args: asyncio.Event().wait()),
+ )
+ await _runtime_insert(blocked, "model", [3], [33])
+ assert not await blocked.close(timeout=0.001)
+
+ asyncio.run(exercise())
+
+
+def test_target_concurrency_burst_is_retained_while_replication_is_blocked() -> None:
+ async def exercise() -> None:
+ replicated = []
+ started = asyncio.Event()
+ release = asyncio.Event()
+
+ async def insert_many(entries):
+ replicated.extend(entries)
+ if len(replicated) == 1:
+ started.set()
+ await release.wait()
+
+ runtime = TokenPrefixRuntime(
+ lookup_shared=lambda _scope, _rendered, _lineage: asyncio.sleep(
+ 0, result=None
+ ),
+ insert_shared_many=insert_many,
+ )
+ await runtime.insert_many([("model-0", [1], [1], "rollout", ())])
+ await started.wait()
+ await asyncio.gather(
+ *(
+ runtime.insert_many(
+ [
+ (
+ f"model-{index}",
+ [index + 1],
+ [index + 1],
+ "rollout",
+ (),
+ )
+ ]
+ )
+ for index in range(1, 32)
+ )
+ )
+ release.set()
+ assert await runtime.flush()
+ assert len(replicated) == 32
+ await runtime.close()
+
+ asyncio.run(exercise())
+
+
+def test_identical_observations_are_idempotent() -> None:
+ cache = TokenPrefixCache()
+ cache.insert([1, 500], [1, 101, 102], "first")
+ cache.insert([1, 500], [1, 101, 102], "second")
+
+ expected = PrefixMatch(2, (1, 101, 102))
+ assert cache.lookup([1, 500, 3], "first") == expected
+ assert cache.lookup([1, 500, 3], "second") == expected
+ assert cache.lookup([1, 500, 3], "unknown") is None
+
+
+def test_eviction_removes_all_variants_at_the_rendered_prefix() -> None:
+ cache = TokenPrefixCache(max_variants=2, max_token_ids=12)
+ cache.insert([1, 500], [1, 101, 102], "shared")
+ cache.insert([1, 500], [1, 500], "shared")
+ cache.insert([9], [9], "third")
+
+ assert cache.lookup([1, 500, 3], "shared") is None
+ assert cache.lookup([9], "third") == PrefixMatch(1, (9,))
+
+
+def test_lineage_presence_tracks_only_retained_observations() -> None:
+ cache = TokenPrefixCache()
+ cache.insert([1], [1], "shared")
+ cache.insert([2], [2], "shared")
+
+ assert cache.lookup([1], "shared") == PrefixMatch(1, (1,))
+ assert cache.evict_oldest() is True
+ assert cache.lookup([1], "shared") == PrefixMatch(1, (1,))
+ assert cache.evict_oldest() is True
+ assert cache.lookup([1], "shared") is None
+
+
+def test_lineage_trimming_is_variant_local() -> None:
+ cache = TokenPrefixCache(max_lineages_per_variant=1)
+ cache.insert([1, 500], [1, 101, 102], "shared")
+ cache.insert([1, 500], [1, 500], "shared")
+ cache.insert([1, 500], [1, 101, 102], "new")
+
+ assert cache.lookup([1, 500, 3], "shared") == PrefixMatch(2, (1, 500))
+ assert cache.lookup([1, 500, 3], "new") == PrefixMatch(2, (1, 101, 102))
+
+
+def test_distinct_lineages_resolve_observed_tokenization_collision() -> None:
+ cache = TokenPrefixCache()
+ cache.insert([1], [1], "first")
+ cache.insert([1, 500], [1, 101, 102], "first")
+ cache.insert([1, 500], [1, 500], "second")
+
+ assert cache.lookup([1, 500, 3], "first") == PrefixMatch(2, (1, 101, 102))
+ assert cache.lookup([1, 500, 3], "second") == PrefixMatch(2, (1, 500))
+
+
+def test_collision_uses_most_recent_variant_at_longest_known_prefix() -> None:
+ cache = TokenPrefixCache()
+ cache.insert([1], [9], "shared")
+ cache.insert([1, 500], [9, 101, 102], "shared")
+ cache.insert([1, 500], [9, 500], "shared")
+
+ assert cache.lookup([1, 500, 3], "shared") == PrefixMatch(2, (9, 500))
+ cache.insert([1, 500], [9, 101, 102], "shared")
+ assert cache.lookup([1, 500, 3], "shared") == PrefixMatch(2, (9, 101, 102))
+
+
+def test_collision_without_shorter_prefix_uses_most_recent_variant() -> None:
+ cache = TokenPrefixCache()
+ cache.insert([1, 500], [1, 101, 102], "shared")
+ cache.insert([1, 500], [1, 500], "shared")
+
+ assert cache.lookup([1, 500, 3], "shared") == PrefixMatch(2, (1, 500))
+
+
+def test_store_partitions_model_keys() -> None:
+ store = TokenPrefixStore()
+ store.insert("base", [1], [2], "lineage")
+ store.insert("checkpoint", [1], [3], "lineage")
+
+ assert store.lookup("base", [1, 4], "lineage") == PrefixMatch(1, (2,))
+ assert store.lookup("checkpoint", [1, 4], "lineage") == PrefixMatch(1, (3,))
+
+
+def test_store_enforces_per_scope_and_global_token_budgets() -> None:
+ store = TokenPrefixStore(
+ max_scopes=4,
+ max_variants_per_scope=4,
+ max_token_ids_per_scope=6,
+ max_token_ids=8,
+ )
+
+ assert store.insert("oversized", [1, 2, 3, 4], [5, 6, 7], "lineage") is False
+ assert store.lookup("oversized", [1, 2, 3, 4], "lineage") is None
+
+ assert store.insert("first", [1, 2], [3, 4], "lineage") is True
+ assert store.insert("second", [5, 6], [7, 8], "lineage") is True
+ assert store.insert("third", [9], [10], "lineage") is True
+
+ assert store.lookup("first", [1, 2], "lineage") is None
+ assert store.lookup("second", [5, 6], "lineage") == PrefixMatch(2, (7, 8))
+ assert store.lookup("third", [9], "lineage") == PrefixMatch(1, (10,))
+
+
+def test_shared_store_reuses_identity_across_engine_replicas() -> None:
+ async def exercise() -> None:
+ shared = TokenPrefixStore()
+
+ async def lookup(scope, rendered, lineage):
+ return shared.lookup(scope, rendered, lineage)
+
+ async def insert(scope, rendered, raw, lineage):
+ shared.insert(scope, rendered, raw, lineage)
+
+ first = TokenPrefixRuntime(
+ lookup_shared=lookup, insert_shared_many=_batch(insert)
+ )
+ second = TokenPrefixRuntime(
+ lookup_shared=lookup, insert_shared_many=_batch(insert)
+ )
+ await _runtime_insert(first, "model", [1, 500], [1, 101, 102])
+ await first.flush()
+
+ canonical, rewritten, _ = await second.rewrite_with_edits(
+ "model",
+ [1, 500, 3],
+ "rollout",
+ shared_candidate=True,
+ )
+
+ assert canonical == [1, 500, 3]
+ assert rewritten == [1, 101, 102, 3]
+
+ asyncio.run(exercise())
+
+
+def test_runtime_migrates_an_implicit_lineage_without_a_client_session() -> None:
+ async def exercise() -> None:
+ shared = TokenPrefixStore()
+ shared.insert("model", [1, 500], [1, 101, 102], "scenario")
+ lookups = []
+
+ async def lookup(scope, rendered, lineage):
+ lookups.append(lineage)
+ return shared.lookup(scope, rendered, lineage)
+
+ async def insert_many(_entries):
+ return None
+
+ runtime = TokenPrefixRuntime(
+ lookup_shared=lookup,
+ insert_shared_many=insert_many,
+ )
+ canonical, rewritten, _ = await runtime.rewrite_with_edits(
+ "model",
+ [1, 500, 3],
+ "tool-call",
+ fallback_lineage="scenario",
+ shared_candidate=True,
+ )
+
+ assert canonical == [1, 500, 3]
+ assert rewritten == [1, 101, 102, 3]
+ assert lookups == ["tool-call", "scenario"]
+ _, local, _ = await runtime.rewrite_with_edits(
+ "model",
+ [1, 500, 4],
+ "tool-call",
+ shared_candidate=False,
+ )
+ assert local == [1, 101, 102, 4]
+ await runtime.close()
+
+ asyncio.run(exercise())
+
+
+def test_shared_store_restores_text_equivalent_noncanonical_sampled_ids() -> None:
+ async def exercise() -> None:
+ shared = TokenPrefixStore()
+
+ async def lookup(scope, rendered, lineage):
+ return shared.lookup(scope, rendered, lineage)
+
+ async def insert(scope, rendered, raw, lineage):
+ shared.insert(scope, rendered, raw, lineage)
+
+ first = TokenPrefixRuntime(
+ lookup_shared=lookup, insert_shared_many=_batch(insert)
+ )
+ second = TokenPrefixRuntime(
+ lookup_shared=lookup, insert_shared_many=_batch(insert)
+ )
+ # Qwen3.5 decodes both forms as " roaming", but only the two-token form
+ # carries the sampled logprobs that ART must train against.
+ await _runtime_insert(first, "model", [10, 65945], [10, 897, 6267])
+ await first.flush()
+
+ canonical, rewritten, edits = await second.rewrite_with_edits(
+ "model", [10, 65945, 20], "rollout", shared_candidate=True
+ )
+
+ assert canonical == [10, 65945, 20]
+ assert rewritten == [10, 897, 6267, 20]
+ assert edits == (PrefixEdit(1, 2, (897, 6267)),)
+
+ asyncio.run(exercise())
+
+
+def test_replication_restarts_when_an_entry_arrives_during_shutdown() -> None:
+ async def exercise() -> None:
+ second = ("model", [3, 4], [3, 5], "rollout", (PrefixEdit(1, 2, (5,)),))
+
+ class RacingQueue(asyncio.Queue):
+ empty_calls = 0
+
+ def empty(self) -> bool:
+ self.empty_calls += 1
+ if self.empty_calls == 2:
+ self.put_nowait([second])
+ return True
+ return super().empty()
+
+ replicated: list[
+ tuple[str, list[int], list[int], str, tuple[PrefixEdit, ...]]
+ ] = []
+
+ async def insert_many(entries):
+ replicated.extend(entries)
+
+ runtime = TokenPrefixRuntime(
+ lookup_shared=lambda *_args: asyncio.sleep(0, result=None),
+ insert_shared_many=insert_many,
+ )
+ setattr(runtime, "_pending", RacingQueue(maxsize=16))
+
+ await _runtime_insert(runtime, "model", [1, 2], [1, 9])
+ assert await runtime.flush(timeout=1)
+
+ assert [entry[1] for entry in replicated] == [[1, 2], [3, 4]]
+ assert await runtime.close()
+
+ asyncio.run(exercise())
+
+
+def test_pending_replication_can_reuse_an_older_matching_shared_prefix() -> None:
+ async def exercise() -> None:
+ started = asyncio.Event()
+ release = asyncio.Event()
+ shared = TokenPrefixStore()
+
+ async def lookup(scope, rendered, lineage):
+ return shared.lookup(scope, rendered, lineage)
+
+ async def insert(scope, rendered, raw, lineage):
+ started.set()
+ await release.wait()
+ shared.insert(scope, rendered, raw, lineage)
+
+ shared.insert("model", [1, 500], [1, 41, 42], "rollout")
+ first = TokenPrefixRuntime(
+ lookup_shared=lookup, insert_shared_many=_batch(insert)
+ )
+ await _runtime_insert(first, "model", [1, 500], [1, 101, 102])
+ await started.wait()
+ second = TokenPrefixRuntime(
+ lookup_shared=lookup, insert_shared_many=_batch(insert)
+ )
+ canonical, rewritten, _ = await second.rewrite_with_edits(
+ "model", [1, 500, 3], "rollout", shared_candidate=True
+ )
+ assert canonical == [1, 500, 3]
+ assert rewritten == [1, 41, 42, 3]
+
+ release.set()
+ assert await first.close()
+ await second.close()
+
+ asyncio.run(exercise())
+
+
+def test_shared_candidate_compares_shared_with_local_identity() -> None:
+ async def exercise() -> None:
+ shared_lookups = 0
+
+ async def lookup(_scope, _rendered, _lineage):
+ nonlocal shared_lookups
+ shared_lookups += 1
+ return None
+
+ runtime = TokenPrefixRuntime(
+ lookup_shared=lookup,
+ insert_shared_many=_discard,
+ )
+ await _runtime_insert(runtime, "model", [1, 2], [1, 2])
+
+ canonical, rewritten, _ = await runtime.rewrite_with_edits(
+ "model", [1, 2, 3], "rollout", shared_candidate=True
+ )
+
+ assert canonical == rewritten == [1, 2, 3]
+ assert shared_lookups == 1
+
+ asyncio.run(exercise())
+
+
+def test_local_nonidentity_and_latest_ambiguity_choice_are_immediate() -> None:
+ async def exercise() -> None:
+ shared_lookups = 0
+
+ async def lookup(_scope, _rendered, _lineage):
+ nonlocal shared_lookups
+ shared_lookups += 1
+ return None
+
+ runtime = TokenPrefixRuntime(
+ lookup_shared=lookup,
+ insert_shared_many=_discard,
+ )
+ await _runtime_insert(runtime, "model", [1, 500], [1, 101, 102])
+ _, first, _ = await runtime.rewrite_with_edits(
+ "model", [1, 500, 3], "rollout", shared_candidate=True
+ )
+ await _runtime_insert(runtime, "model", [1, 500], [1, 500])
+ _, second, _ = await runtime.rewrite_with_edits(
+ "model", [1, 500, 3], "rollout", shared_candidate=True
+ )
+
+ assert first == [1, 101, 102, 3]
+ assert second == [1, 500, 3]
+ assert shared_lookups == 2
+ await runtime.close()
+
+ asyncio.run(exercise())
+
+
+def test_known_lineage_without_a_matching_prefix_consults_shared_lookup() -> None:
+ async def exercise() -> None:
+ shared_lookups = 0
+
+ async def lookup(_scope, _rendered, _lineage):
+ nonlocal shared_lookups
+ shared_lookups += 1
+ return None
+
+ runtime = TokenPrefixRuntime(
+ lookup_shared=lookup,
+ insert_shared_many=_discard,
+ )
+ await _runtime_insert(runtime, "model", [1, 2], [1, 2])
+
+ canonical, rewritten, _ = await runtime.rewrite_with_edits(
+ "model", [8, 9], "rollout", shared_candidate=True
+ )
+
+ assert canonical == rewritten == [8, 9]
+ assert shared_lookups == 1
+
+ asyncio.run(exercise())
+
+
+def test_broad_lineage_without_a_local_match_consults_shared_lookup() -> None:
+ async def exercise() -> None:
+ shared_lookups = 0
+
+ async def lookup(_scope, _rendered, _lineage):
+ nonlocal shared_lookups
+ shared_lookups += 1
+ return PrefixMatch(2, (8, 10))
+
+ runtime = TokenPrefixRuntime(
+ lookup_shared=lookup,
+ insert_shared_many=_discard,
+ )
+ await _runtime_insert(runtime, "model", [1, 2], [1, 2])
+
+ canonical, rewritten, _ = await runtime.rewrite_with_edits(
+ "model",
+ [8, 9, 11],
+ "rollout",
+ shared_candidate=True,
+ )
+
+ assert canonical == [8, 9, 11]
+ assert rewritten == [8, 10, 11]
+ assert shared_lookups == 1
+
+ asyncio.run(exercise())
+
+
+def test_longer_shared_prefix_wins_after_routing_returns_to_stale_worker() -> None:
+ async def exercise() -> None:
+ shared = PrefixMatch(
+ 4,
+ (1, 101, 102, 3, 201, 202),
+ prefix_edits(
+ [1, 500, 3, 600],
+ [1, 101, 102, 3, 201, 202],
+ ),
+ )
+ runtime = TokenPrefixRuntime(
+ lookup_shared=lambda *_args: asyncio.sleep(0, result=shared),
+ insert_shared_many=_discard,
+ )
+ await _runtime_insert(runtime, "model", [1, 500], [1, 101, 102])
+
+ canonical, rewritten, _ = await runtime.rewrite_with_edits(
+ "model",
+ [1, 500, 3, 600, 4],
+ "rollout",
+ shared_candidate=True,
+ )
+
+ assert canonical == [1, 500, 3, 600, 4]
+ assert rewritten == [1, 101, 102, 3, 201, 202, 4]
+ await runtime.close()
+
+ asyncio.run(exercise())
+
+
+def test_unknown_lineage_still_consults_shared_lookup() -> None:
+ async def exercise() -> None:
+ shared_lookups = 0
+
+ async def lookup(_scope, _rendered, _lineage):
+ nonlocal shared_lookups
+ shared_lookups += 1
+ return None
+
+ runtime = TokenPrefixRuntime(
+ lookup_shared=lookup,
+ insert_shared_many=_discard,
+ )
+ await _runtime_insert(runtime, "model", [1, 2], [1, 2], "known")
+ await runtime.rewrite_with_edits(
+ "model",
+ [8, 9],
+ "unknown",
+ shared_candidate=True,
+ )
+
+ assert shared_lookups == 1
+
+ asyncio.run(exercise())
diff --git a/tests/unit/test_tokenizer.py b/tests/unit/test_tokenizer.py
new file mode 100644
index 000000000..a7c8f94a8
--- /dev/null
+++ b/tests/unit/test_tokenizer.py
@@ -0,0 +1,211 @@
+from types import SimpleNamespace
+
+import jinja2
+import pytest
+
+from art import get_tokenizer
+
+
+def test_default_tokenizer_preserves_reasoning_and_honors_explicit_opt_out(monkeypatch):
+ import transformers
+
+ calls = []
+ template = (
+ "{% set enable_thinking = true %}"
+ "{% set ns = namespace(last_query_index=2) %}"
+ "{% for message in messages %}"
+ "{%- if loop.index0 > ns.last_query_index %}{{ message.reasoning }}{% endif %}"
+ "{{ message.content }}{% endfor %}"
+ )
+
+ def load(model, **kwargs):
+ calls.append((model, kwargs))
+ return SimpleNamespace(chat_template=template)
+
+ monkeypatch.setattr(transformers.AutoTokenizer, "from_pretrained", load)
+ tokenizer = get_tokenizer(
+ "model:variant", revision="revision", local_files_only=True
+ )
+ assert calls == [
+ (
+ "model",
+ {
+ "revision": "revision",
+ "local_files_only": True,
+ "trust_remote_code": False,
+ },
+ )
+ ]
+ messages = [{"reasoning": "thought", "content": "answer"}]
+ assert isinstance(tokenizer.chat_template, str)
+ render = jinja2.Environment().from_string(tokenizer.chat_template).render
+ assert render(messages=messages) == "thoughtanswer"
+ assert render(messages=messages, preserve_thinking=False) == "answer"
+ tokenizer.chat_template = "changed"
+ assert get_tokenizer("model").chat_template != "changed"
+
+
+def test_pinned_llama_revision_does_not_switch_repositories(monkeypatch):
+ import transformers
+
+ calls = []
+ monkeypatch.setattr(
+ transformers.AutoTokenizer,
+ "from_pretrained",
+ lambda *a, **k: calls.append((a, k)),
+ )
+ get_tokenizer("meta-llama/Llama-3.1-8B", revision="model-specific-commit")
+ assert calls[0][0] == ("meta-llama/Llama-3.1-8B",)
+
+
+@pytest.mark.parametrize("revision", [None, "pinned-commit"])
+def test_ungated_llama_fallback_is_only_for_unpinned_loads(monkeypatch, revision):
+ import transformers
+
+ calls = []
+
+ def load(model, **kwargs):
+ calls.append(model)
+ if model.startswith("meta-llama/"):
+ raise OSError("gated repository")
+ return SimpleNamespace(chat_template="native text template")
+
+ monkeypatch.setattr(transformers.AutoTokenizer, "from_pretrained", load)
+ if revision:
+ with pytest.raises(OSError, match="gated"):
+ get_tokenizer("meta-llama/Llama-3.1-8B", revision=revision)
+ assert len(calls) == 1
+ else:
+ get_tokenizer("meta-llama/Llama-3.1-8B")
+ assert calls[-1] == "thinkingmachineslabinc/meta-llama-3-instruct-tokenizer"
+
+
+def test_llama_base_retains_native_vocabulary_with_default_chat_template(monkeypatch):
+ import transformers
+
+ native = SimpleNamespace(chat_template=None)
+ monkeypatch.setattr(
+ transformers.AutoTokenizer,
+ "from_pretrained",
+ lambda model, **_: (
+ native
+ if model.startswith("meta-llama/")
+ else SimpleNamespace(chat_template="public chat template")
+ ),
+ )
+ actual = get_tokenizer("meta-llama/Llama-3.1-8B")
+ assert actual is native
+ assert actual.chat_template == "public chat template"
+
+
+@pytest.mark.parametrize(
+ "template",
+ [
+ "reasoning_content and loop.index0 > ns.last_user_index",
+ "thinking_text and loop.index0 > ns_turn.last_user_idx and message.get('tool_calls')",
+ "message.thinking and not future_final_message.found",
+ ],
+)
+def test_supported_template_patches_are_idempotent(template):
+ from art.utils.chat_template import chat_template_with_preserved_thinking
+
+ patched = chat_template_with_preserved_thinking(template)
+ assert patched != template
+ assert chat_template_with_preserved_thinking(patched) == patched
+
+
+def test_glm_preserves_content_whitespace_and_remains_valid_jinja():
+ from art.utils.chat_template import chat_template_with_preserved_thinking
+
+ template = "{% if clear_thinking %}legacy{% endif %}{%- if content.strip() -%}{{ content.strip() }}{%- endif -%}"
+ configured = chat_template_with_preserved_thinking(template)
+ assert isinstance(configured, str)
+ assert (
+ jinja2.Environment().from_string(configured).render(content="\nanswer\n")
+ == "\nanswer\n"
+ )
+ assert chat_template_with_preserved_thinking(configured) == configured
+
+
+def test_kimi_preserves_old_reasoning_when_next_turn_does_not_think():
+ from art.utils.chat_template import chat_template_with_preserved_thinking
+
+ template = """{% set ns = namespace(last_non_tool_call_assistant_msg=0) %}
+ {%- set hist_msgs = messages[:ns.last_non_tool_call_assistant_msg+1] -%}
+ {%- set suffix_msgs = messages[ns.last_non_tool_call_assistant_msg+1:] -%}
+ {% for message in hist_msgs %}{{ message.content }}{% endfor %}
+ {% for message in suffix_msgs %}{%- if thinking is defined and thinking is false -%}{{ message.content }}{% else %}{{ message.reasoning_content }}{{ message.content }}{% endif %}{% endfor %}"""
+ configured = chat_template_with_preserved_thinking(template)
+ assert isinstance(configured, str)
+ render = jinja2.Environment().from_string(configured).render
+ messages = [{"reasoning_content": "thought", "content": "answer"}]
+ assert "thoughtanswer" in render(messages=messages, thinking=False)
+ assert "thought" not in render(
+ messages=messages, thinking=False, preserve_thinking=False
+ )
+ assert chat_template_with_preserved_thinking(configured) == configured
+
+
+@pytest.mark.parametrize("reasoning", [None, "", "\nthought\n"])
+@pytest.mark.parametrize("next_thinking", [False, True])
+@pytest.mark.parametrize(
+ "tools",
+ [
+ None,
+ [
+ {
+ "type": "function",
+ "function": {
+ "name": "lookup",
+ "parameters": {"type": "object", "properties": {}},
+ },
+ }
+ ],
+ ],
+)
+def test_deepseek_v4_retains_prior_turns_when_generation_mode_changes(
+ reasoning, next_thinking, tools
+):
+ from tokenizers import Tokenizer
+ from tokenizers.models import WordLevel
+ from transformers import PreTrainedTokenizerFast
+
+ from art.megatron.dsv4.tokenizer import get_dsv4_tokenizer
+
+ tokenizer = get_dsv4_tokenizer(
+ PreTrainedTokenizerFast(tokenizer_object=Tokenizer(WordLevel({"[UNK]": 0})))
+ )
+ messages = [
+ {"role": "user", "content": "question"},
+ {"role": "assistant", "content": "answer", "reasoning_content": reasoning},
+ ]
+ completed = tokenizer.apply_chat_template(
+ messages, tools=tools, tokenize=False, enable_thinking=reasoning is not None
+ )
+ assert (
+ tokenizer.apply_chat_template(
+ messages,
+ tools=tools,
+ tokenize=False,
+ enable_thinking=reasoning is not None,
+ chat_template=None,
+ )
+ == completed
+ )
+ continued = tokenizer.apply_chat_template(
+ [*messages, {"role": "user", "content": "next"}],
+ tools=tools,
+ tokenize=False,
+ enable_thinking=next_thinking,
+ )
+ assert isinstance(completed, str) and isinstance(continued, str)
+ assert continued.startswith(completed)
+ if reasoning:
+ assert reasoning in continued
+ assert reasoning not in tokenizer.apply_chat_template(
+ [*messages, {"role": "user", "content": "next"}],
+ tools=tools,
+ tokenize=False,
+ enable_thinking=next_thinking,
+ preserve_thinking=False,
+ )
diff --git a/tests/unit/trajectories/test_tokenize.py b/tests/unit/trajectories/test_tokenize.py
index 69f727247..f61777852 100644
--- a/tests/unit/trajectories/test_tokenize.py
+++ b/tests/unit/trajectories/test_tokenize.py
@@ -3421,12 +3421,15 @@ def test_loaded_tokenizers_are_cached_by_model_and_revision(
class AutoTokenizer:
@staticmethod
- def from_pretrained(model: str, *, revision: str | None) -> object:
+ def from_pretrained(
+ model: str, *, revision: str | None = None, **kwargs
+ ) -> object:
loaded.append((model, revision))
return object()
transformers = ModuleType("transformers")
setattr(transformers, "AutoTokenizer", AutoTokenizer)
+ setattr(transformers, "PreTrainedTokenizerFast", AutoTokenizer)
monkeypatch.setitem(sys.modules, "transformers", transformers)
_cached_tokenizer.cache_clear()
try:
@@ -3445,7 +3448,9 @@ def test_deepseek_v4_uses_arts_protocol_renderer(
class AutoTokenizer:
@staticmethod
- def from_pretrained(model: str, *, revision: str | None) -> object:
+ def from_pretrained(
+ model: str, *, revision: str | None = None, **kwargs
+ ) -> object:
assert model == "deepseek-ai/DeepSeek-V4-Flash"
assert revision is None
return raw
@@ -3453,6 +3458,7 @@ def from_pretrained(model: str, *, revision: str | None) -> object:
transformers = ModuleType("transformers")
transformers.__path__ = [] # type: ignore[attr-defined]
setattr(transformers, "AutoTokenizer", AutoTokenizer)
+ setattr(transformers, "PreTrainedTokenizerFast", AutoTokenizer)
tokenizer_base = ModuleType("transformers.tokenization_utils_base")
setattr(tokenizer_base, "PreTrainedTokenizerBase", object)
monkeypatch.setitem(sys.modules, "transformers", transformers)
diff --git a/vllm_runtime/hatch_build.py b/vllm_runtime/hatch_build.py
new file mode 100644
index 000000000..23f31c9f6
--- /dev/null
+++ b/vllm_runtime/hatch_build.py
@@ -0,0 +1,15 @@
+"""Bundle ART's shared sources as real files, including in standalone sdists."""
+
+from pathlib import Path
+
+from hatchling.builders.hooks.plugin.interface import BuildHookInterface
+
+
+class CustomBuildHook(BuildHookInterface):
+ def initialize(self, version, build_data):
+ root = Path(self.root)
+ for source in (root / "src/art_vllm_runtime/_shared").glob("*.py"):
+ if source.is_symlink():
+ build_data["force_include"][str(source.resolve())] = str(
+ source.relative_to(root)
+ )
diff --git a/vllm_runtime/pyproject.toml b/vllm_runtime/pyproject.toml
index 76f8d8d1a..1ed818205 100644
--- a/vllm_runtime/pyproject.toml
+++ b/vllm_runtime/pyproject.toml
@@ -42,8 +42,8 @@ build-backend = "hatchling.build"
[tool.hatch.build.targets.wheel]
packages = ["src/art_vllm_runtime"]
-[tool.hatch.build]
-sources = ["src"]
+[tool.hatch.build.hooks.custom]
+path = "hatch_build.py"
[tool.hatch.metadata]
allow-direct-references = true
diff --git a/vllm_runtime/src/art_vllm_runtime/_shared/__init__.py b/vllm_runtime/src/art_vllm_runtime/_shared/__init__.py
new file mode 100644
index 000000000..e302f59b7
--- /dev/null
+++ b/vllm_runtime/src/art_vllm_runtime/_shared/__init__.py
@@ -0,0 +1 @@
+"""Pure tokenizer/history helpers shared with ART without training imports."""
diff --git a/vllm_runtime/src/art_vllm_runtime/_shared/append_only.py b/vllm_runtime/src/art_vllm_runtime/_shared/append_only.py
new file mode 120000
index 000000000..9b4b33cde
--- /dev/null
+++ b/vllm_runtime/src/art_vllm_runtime/_shared/append_only.py
@@ -0,0 +1 @@
+../../../../src/art_inference/append_only.py
\ No newline at end of file
diff --git a/vllm_runtime/src/art_vllm_runtime/_shared/chat_template.py b/vllm_runtime/src/art_vllm_runtime/_shared/chat_template.py
new file mode 120000
index 000000000..2c3caf7c2
--- /dev/null
+++ b/vllm_runtime/src/art_vllm_runtime/_shared/chat_template.py
@@ -0,0 +1 @@
+../../../../src/art_inference/chat_template.py
\ No newline at end of file
diff --git a/vllm_runtime/src/art_vllm_runtime/_shared/token_prefix.py b/vllm_runtime/src/art_vllm_runtime/_shared/token_prefix.py
new file mode 120000
index 000000000..aa3856092
--- /dev/null
+++ b/vllm_runtime/src/art_vllm_runtime/_shared/token_prefix.py
@@ -0,0 +1 @@
+../../../../src/art_inference/token_prefix.py
\ No newline at end of file
diff --git a/vllm_runtime/src/art_vllm_runtime/_shared/vllm.py b/vllm_runtime/src/art_vllm_runtime/_shared/vllm.py
new file mode 120000
index 000000000..448a69e46
--- /dev/null
+++ b/vllm_runtime/src/art_vllm_runtime/_shared/vllm.py
@@ -0,0 +1 @@
+../../../../src/art_inference/vllm.py
\ No newline at end of file
diff --git a/vllm_runtime/src/art_vllm_runtime/history.py b/vllm_runtime/src/art_vllm_runtime/history.py
new file mode 100644
index 000000000..cdcd3693e
--- /dev/null
+++ b/vllm_runtime/src/art_vllm_runtime/history.py
@@ -0,0 +1,8 @@
+"""ART's shared vLLM history adapter, bundled without training dependencies."""
+
+import sys
+
+from ._shared import vllm as _implementation
+from ._shared.vllm import * # noqa: F403
+
+sys.modules[__name__] = _implementation
diff --git a/vllm_runtime/src/art_vllm_runtime/patches.py b/vllm_runtime/src/art_vllm_runtime/patches.py
index 7c9d9c560..f6a16d97a 100644
--- a/vllm_runtime/src/art_vllm_runtime/patches.py
+++ b/vllm_runtime/src/art_vllm_runtime/patches.py
@@ -9,6 +9,7 @@ def apply_vllm_runtime_patches() -> None:
patch_gemma4_moe_lora_support,
)
from art_vllm_runtime.glm52_patches import apply_glm52_vllm_runtime_patches
+ from art_vllm_runtime.history import patch_history
from art_vllm_runtime.moe_lora_patches import (
patch_local_3d_moe_dummy_lora,
patch_small_batch_moe_lora_intermediate_dtype,
@@ -19,6 +20,7 @@ def apply_vllm_runtime_patches() -> None:
patch_policy_token_spans()
patch_gemma4_moe_lora_support()
subclass_chat_completion_request()
+ patch_history()
patch_nonstreaming_chat_response_offload()
patch_local_3d_moe_dummy_lora()
patch_small_batch_moe_lora_intermediate_dtype()