diff --git a/docs/features/additional-histories.mdx b/docs/features/additional-histories.mdx index 8eddd1669..620bac48e 100644 --- a/docs/features/additional-histories.mdx +++ b/docs/features/additional-histories.mdx @@ -30,7 +30,7 @@ Some models, like Qwen 3, use chat templates that remove special tokens (such as By splitting each turn into a separate history, you can preserve these tokens for training: ```python -from art.trajectories import Trajectory, History +from art.trajectories import LegacyHistory, Trajectory # Instead of a single multi-turn conversation that loses tokens # Train as separate histories to preserve them @@ -41,7 +41,7 @@ trajectory = Trajectory( {"role": "assistant", "content": "I need to add 2 and 24"} ], additional_histories=[ - History( + LegacyHistory( messages_and_choices=[ # The Qwen 3 chat template removes tokens from previous turns {"role": "user", "content": "What is 2+2?"}, @@ -82,7 +82,7 @@ trajectory = Trajectory( ], additional_histories=[ # Sub-agent 1: Code analysis - History( + LegacyHistory( messages_and_choices=[ {"role": "system", "content": "You are a code analysis expert"}, {"role": "user", "content": "Find potential bugs in main.py"}, @@ -90,7 +90,7 @@ trajectory = Trajectory( ] ), # Sub-agent 2: Bug fixing - History( + LegacyHistory( messages_and_choices=[ {"role": "system", "content": "You are a bug fixing expert"}, {"role": "user", "content": "Fix the null pointer issue on line 42"}, @@ -116,7 +116,7 @@ trajectory = Trajectory( ], additional_histories=[ # Previous conversation segment before compaction - History( + LegacyHistory( messages_and_choices=[ {"role": "system", "content": "You are a helpful assistant."}, {"role": "user", "content": "Compacted conversation history: the user asked about quantum entanglement, and the assistant explained..."}, @@ -152,11 +152,11 @@ for history in histories: ### Data Structure -The `History` class structure: +The legacy `LegacyHistory` payload structure: ```python @dataclass -class History: +class LegacyHistory: messages_and_choices: list[dict[str, Any]] tools: list[Tool] | None = None ``` @@ -168,7 +168,7 @@ The `Trajectory` class with additional histories: class Trajectory: messages_and_choices: list[dict[str, Any]] tools: list[Tool] | None = None - additional_histories: list[History] = field(default_factory=list) + additional_histories: list[LegacyHistory] = field(default_factory=list) reward: float | None = None metrics: dict[str, Any] = field(default_factory=dict) ``` @@ -178,7 +178,7 @@ class Trajectory: ### Creating a Trajectory with Additional Histories ```python -from art.trajectories import Trajectory, History +from art.trajectories import LegacyHistory, Trajectory # Create the main conversation main_messages = [ @@ -188,14 +188,14 @@ main_messages = [ ] # Create additional histories -history1 = History( +history1 = LegacyHistory( messages_and_choices=[ {"role": "user", "content": "First subtask"}, {"role": "assistant", "content": "Completing first subtask..."} ] ) -history2 = History( +history2 = LegacyHistory( messages_and_choices=[ {"role": "user", "content": "Second subtask"}, {"role": "assistant", "content": "Completing second subtask..."} diff --git a/pyproject.toml b/pyproject.toml index 20858ccc3..9be997359 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -26,7 +26,7 @@ plotting = ["matplotlib>=3.10.1", "seaborn>=0.13.2"] backend = [ "peft>=0.14.0", "hf-xet>=1.1.0", - "bitsandbytes>=0.45.2", + "bitsandbytes>=0.45.2,!=0.50.0", "unsloth==2026.3.3", "unsloth-zoo==2026.3.1", "torch==2.11.0", diff --git a/src/art/__init__.py b/src/art/__init__.py index f317da1e8..1f8cf15f2 100644 --- a/src/art/__init__.py +++ b/src/art/__init__.py @@ -79,17 +79,28 @@ from .serverless import ServerlessBackend from .trajectories import ( AnthropicMessagesHistory, + AnthropicMessageSource, ChatCompletionsExchange, ChatCompletionsHistory, + ChatCompletionsMessageSource, CompletionsExchange, - CompletionsHistory, + CompletionsSource, + CompletionsStringHistory, + CompletionsStringSourceSpan, + CompletionsTokenHistory, + CompletionsTokenSourceSpan, History, + LegacyHistory, MessagesExchange, ResponsesExchange, ResponsesHistory, + ResponsesItemSource, TokenFlag, + TokenizedHistory, + TokenizedMultiHistoryTrajectory, TokenizedTrajectory, TokenizedTrajectoryGroup, + Tokenizer, Trajectory, TrajectoryExchanges, TrajectoryGroup, @@ -98,10 +109,6 @@ capture_auto_trajectory, # ty: ignore[deprecated] current_trajectory, current_trajectory_group, - tokenize_trajectories, - tokenize_trajectory, - tokenize_trajectory_group, - tokenize_trajectory_groups, trajectory, trajectory_group, ) @@ -158,24 +165,31 @@ "TrajectoryExchanges", "TrajectoryGroup", "History", + "LegacyHistory", "TrajectoryHistory", "ChatCompletionsExchange", "ChatCompletionsHistory", + "ChatCompletionsMessageSource", "CompletionsExchange", - "CompletionsHistory", + "CompletionsSource", + "CompletionsTokenHistory", + "CompletionsStringHistory", + "CompletionsTokenSourceSpan", + "CompletionsStringSourceSpan", "ResponsesExchange", "ResponsesHistory", + "ResponsesItemSource", "MessagesExchange", "AnthropicMessagesHistory", + "AnthropicMessageSource", "TokenFlag", + "Tokenizer", + "TokenizedHistory", + "TokenizedMultiHistoryTrajectory", "TokenizedTrajectory", "TokenizedTrajectoryGroup", "trajectory", "trajectory_group", - "tokenize_trajectory", - "tokenize_trajectories", - "tokenize_trajectory_group", - "tokenize_trajectory_groups", "capture_yielded_trajectory", "yield_trajectory", ] diff --git a/src/art/langgraph/llm_wrapper.py b/src/art/langgraph/llm_wrapper.py index 86e17dc2c..676eaf1f6 100644 --- a/src/art/langgraph/llm_wrapper.py +++ b/src/art/langgraph/llm_wrapper.py @@ -14,7 +14,7 @@ from langchain_core.utils.function_calling import convert_to_openai_tool from langchain_openai import ChatOpenAI -from art.trajectories import History, Trajectory +from art.trajectories import LegacyHistory, Trajectory from .logging import FileLogger from .message_utils import convert_langgraph_messages @@ -89,7 +89,7 @@ def create_messages_from_logs(logger: FileLogger, trajectory: Trajectory): trajectory.tools = tools[idx] else: trajectory.additional_histories.append( - History(messages_and_choices=converted, tools=tools[idx]) + LegacyHistory(messages_and_choices=converted, tools=tools[idx]) ) except Exception: pass diff --git a/src/art/local/backend.py b/src/art/local/backend.py index 64657d0d3..9f93d4832 100644 --- a/src/art/local/backend.py +++ b/src/art/local/backend.py @@ -91,6 +91,7 @@ ) from ..serving_capabilities import ServingCapabilities from ..trajectories import Trajectory, TrajectoryGroup +from ..trajectories._selection import automatic_training_model_selector from ..types import ( Choice, LocalTrainResult, @@ -949,6 +950,12 @@ def _get_packed_tensors( chat_template_tool_schema_format = self._chat_template_tool_schema_format( internal_config ) + model_max_sequence_length = self._model_max_sequence_length(model) + training_max_sequence_length = ( + min(model_max_sequence_length, packed_sequence_length) + if packed_sequence_length is not None + else model_max_sequence_length + ) tokenized_results = list( tokenize_trajectory_groups( tokenizer, @@ -958,11 +965,14 @@ def _get_packed_tensors( image_processor=self._image_processors[model.base_model], chat_template_kwargs=chat_template_kwargs, chat_template_tool_schema_format=chat_template_tool_schema_format, + model=automatic_training_model_selector( + self._model_inference_name(model) + ), + _max_sequence_length=training_max_sequence_length, ) ) if not tokenized_results: return None - model_max_sequence_length = self._model_max_sequence_length(model) too_long_for_model = [ result for result in tokenized_results diff --git a/src/art/openai.py b/src/art/openai.py index 2c99fb044..de3a5e07c 100644 --- a/src/art/openai.py +++ b/src/art/openai.py @@ -163,14 +163,14 @@ def update_chat_completion( tool_call.function.arguments += ( tool_call_delta.function.arguments ) - if getattr(chunk_choice.delta, "reasoning", None): - if not hasattr(choice.message, "reasoning"): - setattr(choice.message, "reasoning", "") + for field in ("reasoning", "reasoning_content"): + value = getattr(chunk_choice.delta, field, None) + if not value: + continue setattr( choice.message, - "reasoning", - getattr(choice.message, "reasoning") - + getattr(chunk_choice.delta, "reasoning"), + field, + (getattr(choice.message, field, None) or "") + value, ) chat_completion.service_tier = chunk.service_tier chat_completion.system_fingerprint = chunk.system_fingerprint diff --git a/src/art/preprocessing/tokenize.py b/src/art/preprocessing/tokenize.py index b466af12a..cca1b2474 100644 --- a/src/art/preprocessing/tokenize.py +++ b/src/art/preprocessing/tokenize.py @@ -1,6 +1,6 @@ from __future__ import annotations -from collections.abc import Callable +from collections.abc import Callable, Iterable from dataclasses import dataclass, field from functools import cached_property from itertools import takewhile @@ -17,14 +17,18 @@ if TYPE_CHECKING: from transformers.image_processing_utils import BaseImageProcessor +from ..openai import ART_MOE_ROUTING_METADATA_KEY from ..trajectories import ( ChatCompletionsExchange, - History, + ChatCompletionsHistory, + ChatCompletionsMessageSource, + LegacyHistory, TokenFlag, Trajectory, TrajectoryGroup, get_messages, ) +from ..trajectories._selection import ModelSelector, resolve_training_model from ..types import MessagesAndChoices from ..utils.chat_template import ( default_chat_template_kwargs_for_tokenizer, @@ -62,6 +66,171 @@ def _flag_spans(flags: list[TokenFlag], flag: TokenFlag) -> list[tuple[int, int] return spans +def _true_spans(mask: list[bool]) -> list[tuple[int, int]]: + return _flag_spans( + [TokenFlag.SAMPLED if value else TokenFlag(0) for value in mask], + TokenFlag.SAMPLED, + ) + + +@dataclass(frozen=True) +class _ChatChoiceTrace: + choices: list[Choice] + offsets: list[int] + lengths: list[int] + + +def _chat_choice_trace( + history: ChatCompletionsHistory, + token_ids: list[int], + flags: list[TokenFlag], +) -> _ChatChoiceTrace | None: + return _chat_source_choice_trace( + ( + (source, source.choice_index) + for source in history.message_sources + if source is not None + ), + token_ids, + flags, + ) + + +def _chat_source_choice_trace( + sources: Iterable[tuple[object, int | None]], + token_ids: list[int], + flags: list[TokenFlag], +) -> _ChatChoiceTrace | None: + """Recover per-choice boundaries that adjacent SAMPLED flags cannot represent.""" + + sourced_choices: list[Choice] = [] + seen: set[tuple[int, int]] = set() + for source, choice_index in sources: + exchange = ( + source.exchange + if isinstance(source, ChatCompletionsMessageSource) + else source + ) + if not isinstance(exchange, ChatCompletionsExchange) or choice_index is None: + continue + key = (id(exchange), choice_index) + if key in seen: + continue + seen.add(key) + sourced_choices.append( + next( + item for item in exchange.response.choices if item.index == choice_index + ) + ) + if not sourced_choices: + return None + has_routing = any( + choice_moe_routing_metadata(choice) is not None for choice in sourced_choices + ) + if any(choice_vllm_token_metadata(choice) is None for choice in sourced_choices): + if has_routing: + raise RuntimeError( + "MoE routing replay requires exact token IDs for every sourced choice" + ) + return None + + metadata = [choice_vllm_token_metadata(choice) for choice in sourced_choices] + if any(value is None for value in metadata): + raise AssertionError("choice metadata was checked above") + exact_metadata = cast(list[tuple[list[int], list[int]]], metadata) + + choices: list[Choice] = [] + offsets: list[int] = [] + lengths: list[int] = [] + cursor = 0 + for position, (choice, (prompt_ids, completion_ids)) in enumerate( + zip(sourced_choices, exact_metadata, strict=True) + ): + offset = len(prompt_ids) + if ( + offset < cursor + or offset > len(token_ids) + or token_ids[:offset] != prompt_ids + ): + if has_routing: + raise RuntimeError( + "MoE routed prompt tokens are absent from tokenized history" + ) + return None + next_boundary = ( + len(exact_metadata[position + 1][0]) + if position + 1 < len(exact_metadata) + else len(token_ids) + ) + end = offset + while ( + end < min(next_boundary, len(token_ids)) and flags[end] & TokenFlag.SAMPLED + ): + end += 1 + retained_ids = token_ids[offset:end] + if not retained_ids or not completion_ids: + if has_routing: + raise RuntimeError( + "MoE routed completion tokens are absent from tokenized history" + ) + return None + retained_start = len(completion_ids) - len(retained_ids) + if retained_start < 0 or completion_ids[retained_start:] != retained_ids: + if has_routing: + raise RuntimeError( + "MoE routed completion suffix disagrees with tokenized history" + ) + return None + choices.append( + _choice_with_retained_routing(choice, retained_start) + if retained_start + else choice + ) + offsets.append(offset) + lengths.append(len(retained_ids)) + cursor = offset + len(retained_ids) + return ( + _ChatChoiceTrace(choices=choices, offsets=offsets, lengths=lengths) + if choices + else None + ) + + +def _choice_with_retained_routing(choice: Choice, start: int) -> Choice: + """Copy one routed choice while retaining only a completion suffix.""" + + routing = choice_moe_routing_metadata(choice) + if routing is None: + return choice + token_metadata = choice_vllm_token_metadata(choice) + assert token_metadata is not None + prompt_ids, completion_ids = token_metadata + routes = routing.get("routed_experts") + if not isinstance(routes, np.ndarray): + raise RuntimeError("Missing binary routed experts") + completion_route_count = len(routes) - len(prompt_ids) + if completion_route_count not in { + len(completion_ids), + max(len(completion_ids) - 1, 0), + }: + raise RuntimeError( + "routed_experts length does not match prompt/completion token ids" + ) + retained = choice.model_copy() + extra = retained.model_extra + if extra is None: + raise RuntimeError("OpenAI Choice.model_extra is unavailable for route replay") + extra["token_ids"] = completion_ids[start:] + extra[ART_MOE_ROUTING_METADATA_KEY] = { + **routing, + "completion_token_ids": completion_ids[start:], + "routed_experts": np.concatenate( + (routes[: len(prompt_ids)], routes[len(prompt_ids) + start :]) + ), + } + return retained + + def _slice_moe_routes( routes: MoeRouteArray | MoeRouteSegments | None, start: int ) -> MoeRouteArray | MoeRouteSegments | None: @@ -448,7 +617,7 @@ def _tokenized_result_from_vllm_choices( def assemble_vllm_training_sequences( *, tokenizer: PreTrainedTokenizerBase, - histories: list[History], + histories: list[LegacyHistory], advantage: float, allow_training_without_logprobs: bool, trajectory: Trajectory, @@ -548,6 +717,8 @@ def tokenize_trajectory_groups( image_processor: BaseImageProcessor | None = None, chat_template_kwargs: dict[str, Any] | None = None, chat_template_tool_schema_format: ChatTemplateToolSchemaFormat = "default", + model: ModelSelector | str | None = None, + _max_sequence_length: int | None = None, ) -> Generator["TokenizedResult", None, None]: for group in trajectory_groups: if not group: @@ -569,92 +740,142 @@ def tokenize_trajectory_groups( if trajectory.exchanges: from ..trajectories._tokenize import ( _as_tokenizer, - _exchange_list, - _exchange_tokens, - tokenize_one, + _first_introduction_mask, + _require_causal_predecessor, + _SampledSourceKey, + _tokenize_trajectory_with_trace, ) - exchange_result = tokenize_one( + selected_model = resolve_training_model(trajectory, model) + exchange_results, traces = _tokenize_trajectory_with_trace( trajectory, - tokenizer.name_or_path, - model=None, + model=selected_model, + base_model=tokenizer.name_or_path, + tokenizer=_as_tokenizer(tokenizer), chat_template=None, chat_template_kwargs=chat_template_kwargs, - tokenizer_instance=_as_tokenizer(tokenizer), ) - sampled = [ - bool(flag & TokenFlag.SAMPLED) for flag in exchange_result.flags - ] - sampled_spans = _flag_spans(exchange_result.flags, TokenFlag.SAMPLED) - choice_spans = sampled_spans - if not allow_training_without_logprobs and any( - trainable and math.isnan(logprob) - for trainable, logprob in zip( - sampled, - exchange_result.logprobs, - strict=True, - ) + trajectory_results = [] + seen_source_keys: set[_SampledSourceKey] = set() + for exchange_result, trace in zip( + exchange_results.histories, traces, strict=True ): - raise RuntimeError( - "Exchange trajectory is missing logprobs for trainable tokens" - ) - exchanges = _exchange_list(trajectory, None) - chat_choices = [ - exchange.response.choices[0] - for exchange in exchanges - if isinstance(exchange, ChatCompletionsExchange) - ] - if len(chat_choices) == len(exchanges): - exact_spans: list[tuple[int, int]] = [] - for exchange in exchanges: - prompt, completion, _ = _exchange_tokens(exchange) - if prompt is None or completion is None: - exact_spans = [] - break - exact_spans.append((len(prompt), len(prompt) + len(completion))) - choice_spans = exact_spans or choice_spans - moe_routes, moe_stats = align_choice_routes_to_tokenized_result( - token_ids=exchange_result.token_ids, - choices=chat_choices, - choice_offsets=[start for start, _ in choice_spans], - choice_token_lengths=[ - end - start for start, end in choice_spans - ], + if ( + _max_sequence_length is not None + and len(exchange_result.token_ids) > _max_sequence_length + ): + preview_seen = set(seen_source_keys) + would_train = _first_introduction_mask( + trace.source_keys, preview_seen + ) + if any(would_train): + trajectory_results.append( + TokenizedResult( + advantage=advantage, + chat="", + token_ids=exchange_result.token_ids, + input_pos=list( + range(len(exchange_result.token_ids)) + ), + assistant_mask=[0] * len(exchange_result.token_ids), + logprobs=exchange_result.logprobs, + pixel_values=None, + image_grid_thw=None, + trajectory=trajectory, + choice_offsets=[], + extra_logprobs={}, + moe_routed_experts=None, + moe_routing_alignment_stats=MoeRoutingAlignmentStats(), + _tokenizer=tokenizer, + ) + ) + continue + trainable = _first_introduction_mask( + trace.source_keys, seen_source_keys ) - else: - if any( - choice_moe_routing_metadata(choice) is not None - for choice in chat_choices + _require_causal_predecessor(trainable) + if not any(trainable): + continue + choice_spans = _true_spans(trainable) + if not allow_training_without_logprobs and any( + selected and math.isnan(logprob) + for selected, logprob in zip( + trainable, + exchange_result.logprobs, + strict=True, + ) ): raise RuntimeError( - "MoE routing replay requires an all-Chat-Completions " - "exchange trajectory" + "Exchange trajectory is missing logprobs for trainable tokens" + ) + ordered_source_keys = dict.fromkeys( + key for key in trace.source_keys if key is not None + ) + chat_trace = _chat_source_choice_trace( + ( + (trace.sources[key], key.index) + for key in ordered_source_keys + ), + exchange_result.token_ids, + exchange_result.flags, + ) + if chat_trace is not None: + selected_choices = [ + index + for index, (start, length) in enumerate( + zip( + chat_trace.offsets, + chat_trace.lengths, + strict=True, + ) + ) + if any(trainable[start : start + length]) + ] + moe_routes, moe_stats = align_choice_routes_to_tokenized_result( + token_ids=exchange_result.token_ids, + choices=[ + chat_trace.choices[index] for index in selected_choices + ], + choice_offsets=[ + chat_trace.offsets[index] for index in selected_choices + ], + choice_token_lengths=[ + chat_trace.lengths[index] for index in selected_choices + ], + ) + choice_spans = [ + ( + chat_trace.offsets[index], + chat_trace.offsets[index] + chat_trace.lengths[index], + ) + for index in selected_choices + ] + else: + moe_routes = None + moe_stats = MoeRoutingAlignmentStats() + trajectory_results.append( + TokenizedResult( + advantage=advantage, + chat="", + token_ids=exchange_result.token_ids, + input_pos=list(range(len(exchange_result.token_ids))), + assistant_mask=[int(value) for value in trainable], + logprobs=exchange_result.logprobs, + pixel_values=None, + image_grid_thw=None, + trajectory=trajectory, + choice_offsets=[start for start, _ in choice_spans], + extra_logprobs={}, + moe_routed_experts=moe_routes, + moe_routing_alignment_stats=moe_stats, + _tokenizer=tokenizer, ) - moe_routes = None - moe_stats = MoeRoutingAlignmentStats() - trajectory_results = [ - TokenizedResult( - advantage=advantage, - chat="", - token_ids=exchange_result.token_ids, - input_pos=list(range(len(exchange_result.token_ids))), - assistant_mask=[int(value) for value in sampled], - logprobs=exchange_result.logprobs, - pixel_values=None, - image_grid_thw=None, - trajectory=trajectory, - choice_offsets=[start for start, _ in choice_spans], - extra_logprobs={}, - moe_routed_experts=moe_routes, - moe_routing_alignment_stats=moe_stats, - _tokenizer=tokenizer, ) - ] else: trajectory_results = assemble_vllm_training_sequences( tokenizer=tokenizer, histories=[ - History( + LegacyHistory( messages_and_choices=trajectory.messages_and_choices, tools=trajectory.tools, ), @@ -707,7 +928,7 @@ def tokenize_trajectory_groups( def tokenize_trajectory( tokenizer: "PreTrainedTokenizerBase", image_processor: BaseImageProcessor | None, - history: History, + history: LegacyHistory, advantage: float, allow_training_without_logprobs: bool, trajectory: Trajectory, diff --git a/src/art/tau_bench/rollout.py b/src/art/tau_bench/rollout.py index 093bd6396..8e058428f 100644 --- a/src/art/tau_bench/rollout.py +++ b/src/art/tau_bench/rollout.py @@ -3,9 +3,10 @@ from collections.abc import Mapping import json import os -from typing import Any, overload +from typing import Any, cast, overload from openai import AsyncOpenAI, BadRequestError +from openai.types.chat import ChatCompletionMessageParam from openai.types.completion_usage import CompletionUsage from art.costs import get_model_pricing, tokens_to_cost @@ -94,12 +95,12 @@ async def rollout( model=model, base_model=base_model, ) + messages: list[ChatCompletionMessageParam] = [ + {"role": "system", "content": env.info["policy"]}, + {"role": "user", "content": env.observation.removeprefix("user: ")}, + ] + tools = env.info.get("tools") or [] trajectory = Trajectory( - messages_and_choices=[ - {"role": "system", "content": env.info["policy"]}, - {"role": "user", "content": env.observation.removeprefix("user: ")}, - ], - tools=env.info.get("tools"), reward=0, metrics={ "cost/tinker/prefill": 0.0, @@ -110,65 +111,76 @@ async def rollout( ) terminated = False num_turns = 0 - while not terminated: - if max_turns is not None and num_turns >= max_turns: - break - try: - chat_completion = await openai_client.chat.completions.create( - messages=trajectory.messages(), - model=model_name, - stream=False, - tool_choice="auto", - tools=trajectory.tools or [], - **chat_completion_kwargs, - ) - except BadRequestError as exc: - if _is_max_tokens_error(exc): + with trajectory: + while not terminated: + if max_turns is not None and num_turns >= max_turns: break - raise - _record_tinker_costs( - trajectory, - cost_model, - chat_completion.usage, - assert_costs=assert_costs, - ) - choice = chat_completion.choices[0] - trajectory.messages_and_choices.append(choice) - tool_calls = getattr(choice.message, "tool_calls", None) - if tool_calls: - for tool_call in tool_calls: - action = _tool_call_action(tool_call) - step = await client.step_environment(env.id, action) - trajectory.messages_and_choices.append( + try: + chat_completion = await openai_client.chat.completions.create( + messages=messages, + model=model_name, + stream=False, + tool_choice="auto", + tools=tools, + **chat_completion_kwargs, + ) + except BadRequestError as exc: + if _is_max_tokens_error(exc): + break + raise + _record_tinker_costs( + trajectory, + cost_model, + chat_completion.usage, + assert_costs=assert_costs, + ) + choice = chat_completion.choices[0] + messages.append( + cast( + ChatCompletionMessageParam, + choice.message.model_dump(exclude_none=True), + ) + ) + tool_calls = getattr(choice.message, "tool_calls", None) + if tool_calls: + for tool_call in tool_calls: + action = _tool_call_action(tool_call) + step = await client.step_environment(env.id, action) + messages.append( + { + "role": "tool", + "content": step.observation.removeprefix("tool: "), + "tool_call_id": tool_call.id, + } + ) + trajectory.reward += step.reward + terminated = step.terminated + else: + step = await client.step_environment( + env.id, + choice.message.content or "", + ) + if "user_message_cost" in step.info: + trajectory.metrics["cost/user"] += step.info[ + "user_message_cost" + ] + elif assert_costs: + raise ValueError("Costs are not supported for the user model") + messages.append( { - "role": "tool", - "content": step.observation.removeprefix("tool: "), - "tool_call_id": tool_call.id, + "role": "user", + "content": step.observation.removeprefix("user: "), } ) trajectory.reward += step.reward terminated = step.terminated - else: - step = await client.step_environment( - env.id, - choice.message.content or "", - ) - if "user_message_cost" in step.info: - trajectory.metrics["cost/user"] += step.info["user_message_cost"] - elif assert_costs: - raise ValueError("Costs are not supported for the user model") - trajectory.messages_and_choices.append( - {"role": "user", "content": step.observation.removeprefix("user: ")} - ) - trajectory.reward += step.reward - terminated = step.terminated - num_turns += 1 - usage = chat_completion.usage - if usage is not None and _would_exceed_context_limit( - usage.total_tokens, - _requested_completion_tokens(chat_completion_kwargs), - ): - break + num_turns += 1 + usage = chat_completion.usage + if usage is not None and _would_exceed_context_limit( + usage.total_tokens, + _requested_completion_tokens(chat_completion_kwargs), + ): + break trajectory.metrics["num_turns"] = num_turns return trajectory diff --git a/src/art/tinker/server.py b/src/art/tinker/server.py index 7abe85865..1767fa868 100644 --- a/src/art/tinker/server.py +++ b/src/art/tinker/server.py @@ -572,12 +572,12 @@ async def messages_and_choices_prompt_tokens_and_choice_offsets( tools: Tools | None, ) -> tuple[list[int], list[int]] | None: from art.preprocessing.tokenize import tokenize_trajectory - from art.trajectories import History, Trajectory + from art.trajectories import LegacyHistory, Trajectory result = tokenize_trajectory( tokenizer=self._get_renderer(base_model).tokenizer, image_processor=None, - history=History( + history=LegacyHistory( messages_and_choices=messages_and_choices, tools=tools, ), diff --git a/src/art/tinker_native/backend.py b/src/art/tinker_native/backend.py index d397adc85..2aba95947 100644 --- a/src/art/tinker_native/backend.py +++ b/src/art/tinker_native/backend.py @@ -41,6 +41,7 @@ from ..tinker.backend import get_renderer_name from ..tinker.server import get_free_port from ..trajectories import Trajectory, TrajectoryGroup +from ..trajectories._selection import automatic_training_model_selector from ..types import TrainResult, TrainSFTConfig from ..utils.lifecycle import process_shutdown_timeout from ..utils.output_dirs import get_model_dir @@ -358,6 +359,7 @@ async def train( state.tokenizer, normalize_advantages, base_model=model.base_model, + model=automatic_training_model_selector(self._model_inference_name(model)), ) metrics: dict[str, float] = { diff --git a/src/art/tinker_native/data.py b/src/art/tinker_native/data.py index 4fd14e8ef..e0051e57f 100644 --- a/src/art/tinker_native/data.py +++ b/src/art/tinker_native/data.py @@ -10,14 +10,14 @@ import torch from ..trajectories import ( - History, + LegacyHistory, TokenFlag, - TokenizedTrajectory, + TokenizedHistory, Trajectory, TrajectoryGroup, get_messages, - tokenize_trajectory, ) +from ..trajectories._selection import ModelSelector, resolve_training_model from ..types import MessagesAndChoices @@ -143,6 +143,7 @@ def trajectory_groups_to_datums( normalize_advantages: bool = True, *, base_model: str | None = None, + model: ModelSelector | str | None = None, ) -> list[tinker.Datum]: datums: list[tinker.Datum] = [] @@ -163,11 +164,32 @@ def trajectory_groups_to_datums( continue for trajectory, advantage in zip(group.trajectories, advantages): if trajectory.exchanges: - datum = _tokenized_trajectory_to_datum( - tokenize_trajectory(trajectory, base_model=base_model), advantage + from ..trajectories._tokenize import ( + _as_tokenizer, + _first_introduction_mask, + _SampledSourceKey, + _tokenize_trajectory_with_trace, ) - if datum is not None: - datums.append(datum) + + selected_model = resolve_training_model(trajectory, model) + tokenized, traces = _tokenize_trajectory_with_trace( + trajectory, + model=selected_model, + base_model=base_model, + tokenizer=_as_tokenizer(tokenizer) + if tokenizer is not None + else None, + ) + seen_source_keys: set[_SampledSourceKey] = set() + for history, trace in zip(tokenized.histories, traces, strict=True): + trainable = _first_introduction_mask( + trace.source_keys, seen_source_keys + ) + datum = _tokenized_trajectory_to_datum( + history, advantage, trainable=trainable + ) + if datum is not None: + datums.append(datum) continue for history in iter_trajectory_histories(trajectory): datum = history_to_datum(history, advantage, renderer, tokenizer) @@ -177,8 +199,8 @@ def trajectory_groups_to_datums( return datums -def iter_trajectory_histories(trajectory: Trajectory) -> Iterable[History]: - yield History( +def iter_trajectory_histories(trajectory: Trajectory) -> Iterable[LegacyHistory]: + yield LegacyHistory( messages_and_choices=trajectory.messages_and_choices, tools=trajectory.tools, ) @@ -222,7 +244,7 @@ def extract_logprobs_from_choice( def history_to_datum( - history: History, + history: LegacyHistory, advantage: float, renderer: Any, tokenizer: Any, @@ -283,17 +305,31 @@ def build_datum( def _tokenized_trajectory_to_datum( - tokenized: TokenizedTrajectory, advantage: float + tokenized: TokenizedHistory, + advantage: float, + *, + trainable: list[bool] | None = None, ) -> tinker.Datum | None: - sampled = [bool(flag & TokenFlag.SAMPLED) for flag in tokenized.flags] - if not (len(tokenized.token_ids) == len(tokenized.logprobs) == len(sampled)): + if trainable is None: + trainable = [bool(flag & TokenFlag.SAMPLED) for flag in tokenized.flags] + if not ( + len(tokenized.token_ids) + == len(tokenized.logprobs) + == len(tokenized.flags) + == len(trainable) + ): raise ValueError("Tokenized trajectory fields differ in length") - if len(tokenized.token_ids) < 2 or not any(sampled): - return None - if sampled[0]: + if any( + selected and not flag & TokenFlag.SAMPLED + for selected, flag in zip(trainable, tokenized.flags, strict=True) + ): + raise ValueError("Only sampled tokens can be selected for Tinker training") + if trainable and trainable[0]: raise ValueError("A trainable trajectory cannot start with a sampled token") + if len(tokenized.token_ids) < 2 or not any(trainable): + return None - action_mask = sampled[1:] + action_mask = trainable[1:] if any( trainable and math.isnan(logprob) for trainable, logprob in zip(action_mask, tokenized.logprobs[1:], strict=True) diff --git a/src/art/trajectories/__init__.py b/src/art/trajectories/__init__.py index fa0c9a87a..0e1068cf9 100644 --- a/src/art/trajectories/__init__.py +++ b/src/art/trajectories/__init__.py @@ -9,11 +9,21 @@ Mapping, ) from contextlib import asynccontextmanager +from dataclasses import dataclass from datetime import datetime from enum import IntFlag import time from types import TracebackType -from typing import Annotated, Any, Literal, TypeAlias, overload +from typing import ( + Annotated, + Any, + Generic, + Literal, + Protocol, + TypeAlias, + TypeVar, + overload, +) from anthropic.types import ( Message as AnthropicMessage, @@ -43,6 +53,9 @@ from openai.types.responses import ( ToolParam as ResponsesToolParam, ) +from openai.types.responses.response_create_params import ( + Conversation as ResponsesConversation, +) import pydantic from typing_extensions import TypedDict, deprecated @@ -57,6 +70,23 @@ MetadataValue = Any +class Tokenizer(Protocol): + """Minimal tokenizer surface used by trajectory tokenization.""" + + def __call__(self, text: str, *, add_special_tokens: bool = False) -> object: ... + + def apply_chat_template( + self, + messages: list[dict[str, Any]], + *, + tools: object, + tokenize: bool, + add_generation_prompt: bool, + chat_template: str | None = None, + **kwargs: object, + ) -> object: ... + + class TokenFlag(IntFlag): """Independent facts about a token; members may be combined.""" @@ -96,6 +126,7 @@ class CompletionsRequest(TypedDict, total=False, extra_items=Any): echo: bool stop: str | list[str] seed: int + suffix: str class ResponsesRequest(TypedDict, total=False, extra_items=Any): @@ -105,6 +136,7 @@ class ResponsesRequest(TypedDict, total=False, extra_items=Any): input: str | ResponseInputParam instructions: str previous_response_id: str + conversation: ResponsesConversation stream: bool tools: list[ResponsesToolParam] max_output_tokens: int @@ -216,7 +248,7 @@ class PydanticException(pydantic.BaseModel): traceback: str -class History(pydantic.BaseModel): +class LegacyHistory(pydantic.BaseModel): messages_and_choices: MessagesAndChoices tools: Tools | None = None @@ -227,18 +259,114 @@ def serialize_messages_and_choices(self, value: MessagesAndChoices) -> list[Any] def messages(self) -> Messages: return get_messages(self.messages_and_choices) - def as_chat_completions_history(self) -> ChatCompletionsHistory: + def as_chat_completions_history( + self, *, model: str | None = None + ) -> ChatCompletionsHistory: from ._history import legacy_as_chat_completions_history - return legacy_as_chat_completions_history(self) + return legacy_as_chat_completions_history(self, model=model) + + def tokenize( + self, + *, + model: str, + base_model: str | None = None, + tokenizer: Tokenizer | None = None, + chat_template: str | None = None, + chat_template_kwargs: Mapping[str, object] | None = None, + ) -> TokenizedHistory: + from ._tokenize import tokenize_history + + return tokenize_history( + self, + model=model, + base_model=base_model, + tokenizer=tokenizer, + chat_template=chat_template, + chat_template_kwargs=chat_template_kwargs, + ) -class ChatCompletionsHistory(pydantic.BaseModel): +class History: + """Mutable, protocol-native view of one tokenizable sequence.""" + model: str | None - messages: Annotated[pydantic.SerializeAsAny[Messages], pydantic.SkipValidation] - tools: Annotated[pydantic.SerializeAsAny[Tools | None], pydantic.SkipValidation] = ( - None - ) + + def as_chat_completions_history(self) -> ChatCompletionsHistory: + raise ValueError( + f"{type(self).__name__} cannot be represented as Chat Completions" + ) + + def tokenize( + self, + *, + base_model: str | None = None, + tokenizer: Tokenizer | None = None, + chat_template: str | None = None, + chat_template_kwargs: Mapping[str, object] | None = None, + ) -> TokenizedHistory: + from ._tokenize import tokenize_history + + return tokenize_history( + self, + model=self.model, + base_model=base_model, + tokenizer=tokenizer, + chat_template=chat_template, + chat_template_kwargs=chat_template_kwargs, + ) + + +@dataclass(frozen=True, slots=True) +class ChatCompletionsMessageSource: + exchange: ChatCompletionsExchange | MessagesExchange | ResponsesExchange + request_index: int | None = None + choice_index: int | None = None + output_indices: tuple[int, ...] | None = None + generation_index: int | None = None + + +@dataclass(frozen=True, slots=True) +class AnthropicMessageSource: + exchange: MessagesExchange + request_index: int | None = None + + +@dataclass(frozen=True, slots=True) +class ResponsesItemSource: + exchange: ResponsesExchange + request_index: int | None = None + output_index: int | None = None + generation_index: int | None = None + + +@dataclass(frozen=True, slots=True) +class CompletionsSource: + exchange: CompletionsExchange + prompt_index: int + choice_index: int | None = None + + +@dataclass(frozen=True, slots=True) +class CompletionsTokenSourceSpan: + start: int + end: int + source: CompletionsSource | None + + +@dataclass(frozen=True, slots=True) +class CompletionsStringSourceSpan: + start: int + end: int + source: CompletionsSource | None + + +@dataclass +class ChatCompletionsHistory(History): + model: str | None + messages: Messages + message_sources: list[ChatCompletionsMessageSource | None] + tools: Tools | None = None chat_template: str | None = None chat_template_kwargs: dict[str, Any] | None = None @@ -246,23 +374,14 @@ def as_chat_completions_history(self) -> ChatCompletionsHistory: return self -class AnthropicMessagesHistory(pydantic.BaseModel): +@dataclass +class AnthropicMessagesHistory(History): model: str - system: Annotated[ - pydantic.SerializeAsAny[str | list[AnthropicTextBlockParam] | None], - pydantic.SkipValidation, - ] = None - messages: Annotated[ - pydantic.SerializeAsAny[list[AnthropicMessageParam]], pydantic.SkipValidation - ] - tools: Annotated[ - pydantic.SerializeAsAny[list[AnthropicToolParam] | None], - pydantic.SkipValidation, - ] = None - thinking: Annotated[ - pydantic.SerializeAsAny[AnthropicThinkingConfigParam | None], - pydantic.SkipValidation, - ] = None + messages: list[AnthropicMessageParam] + message_sources: list[AnthropicMessageSource | None] + system: str | list[AnthropicTextBlockParam] | None = None + system_source: MessagesExchange | None = None + tools: list[AnthropicToolParam] | None = None chat_template: str | None = None chat_template_kwargs: dict[str, Any] | None = None @@ -272,16 +391,16 @@ def as_chat_completions_history(self) -> ChatCompletionsHistory: return anthropic_as_chat_completions_history(self) -class ResponsesHistory(pydantic.BaseModel): +@dataclass +class ResponsesHistory(History): model: str - input: Annotated[ - pydantic.SerializeAsAny[ResponseInputParam], pydantic.SkipValidation - ] + input: ResponseInputParam + input_sources: list[ResponsesItemSource | None] instructions: str | None = None - tools: Annotated[ - pydantic.SerializeAsAny[list[ResponsesToolParam] | None], - pydantic.SkipValidation, - ] = None + instructions_source: ResponsesExchange | None = None + tools: list[ResponsesToolParam] | None = None + conversation: ResponsesConversation | None = None + previous_response_id: str | None = None chat_template: str | None = None chat_template_kwargs: dict[str, Any] | None = None @@ -291,9 +410,22 @@ def as_chat_completions_history(self) -> ChatCompletionsHistory: return responses_as_chat_completions_history(self) -class CompletionsHistory(pydantic.BaseModel): +@dataclass +class CompletionsTokenHistory(History): model: str - token_ids: list[int] + prompt: list[int] + prompt_sources: list[CompletionsTokenSourceSpan] + sampled_spans: list[tuple[int, int]] + + def as_chat_completions_history(self) -> ChatCompletionsHistory: + raise ValueError("Raw Completions history has no chat-message structure") + + +@dataclass +class CompletionsStringHistory(History): + model: str + prompt: str + prompt_sources: list[CompletionsStringSourceSpan] sampled_spans: list[tuple[int, int]] def as_chat_completions_history(self) -> ChatCompletionsHistory: @@ -301,11 +433,12 @@ def as_chat_completions_history(self) -> ChatCompletionsHistory: TrajectoryHistory: TypeAlias = ( - History + LegacyHistory | ChatCompletionsHistory | AnthropicMessagesHistory | ResponsesHistory - | CompletionsHistory + | CompletionsTokenHistory + | CompletionsStringHistory ) @@ -316,7 +449,7 @@ class Trajectory(_CompactModel): exclude_if=lambda value: not value, ) tools: Tools | None = None - additional_histories: list[History] = pydantic.Field( + additional_histories: list[LegacyHistory] = pydantic.Field( default_factory=list, ) reward: float = 0.0 @@ -378,6 +511,7 @@ async def track_duration(self, metric_name: str) -> AsyncGenerator[None, None]: def __str__(self) -> str: return f"Trajectory(reward={self.reward}, metrics={self.metrics}, metadata={self.metadata})" + # Every model selector accepts an exact identity or a shell-style pattern. def chat_completions_history( self, *, model: str | None = None ) -> ChatCompletionsHistory: @@ -385,6 +519,13 @@ def chat_completions_history( return chat_completions_history(self, model=model) + def chat_completions_histories( + self, *, model: str | None = None + ) -> list[ChatCompletionsHistory]: + from ._history import chat_completions_histories + + return chat_completions_histories(self, model=model) + def anthropic_messages_history( self, *, model: str | None = None ) -> AnthropicMessagesHistory: @@ -392,21 +533,109 @@ def anthropic_messages_history( return anthropic_messages_history(self, model=model) + def anthropic_messages_histories( + self, *, model: str | None = None + ) -> list[AnthropicMessagesHistory]: + from ._history import anthropic_messages_histories + + return anthropic_messages_histories(self, model=model) + def responses_history(self, *, model: str | None = None) -> ResponsesHistory: from ._history import responses_history return responses_history(self, model=model) - def completions_history(self, *, model: str | None = None) -> CompletionsHistory: - from ._history import completions_history + def responses_histories( + self, *, model: str | None = None + ) -> list[ResponsesHistory]: + from ._history import responses_histories - return completions_history(self, model=model) + return responses_histories(self, model=model) + + def completions_token_history( + self, *, model: str | None = None + ) -> CompletionsTokenHistory: + from ._history import completions_token_history + + return completions_token_history(self, model=model) + + def completions_token_histories( + self, *, model: str | None = None + ) -> list[CompletionsTokenHistory]: + from ._history import completions_token_histories + + return completions_token_histories(self, model=model) + + def completions_string_history( + self, *, model: str | None = None + ) -> CompletionsStringHistory: + from ._history import completions_string_history + + return completions_string_history(self, model=model) + + def completions_string_histories( + self, *, model: str | None = None + ) -> list[CompletionsStringHistory]: + from ._history import completions_string_histories + + return completions_string_histories(self, model=model) def history(self, *, model: str | None = None) -> TrajectoryHistory: from ._history import trajectory_history return trajectory_history(self, model=model) + def histories(self, *, model: str | None = None) -> list[TrajectoryHistory]: + from ._history import trajectory_histories + + return trajectory_histories(self, model=model) + + @overload + def tokenize( + self, + *, + multi_history: Literal[False] = False, + model: str | None = None, + base_model: str | None = None, + tokenizer: Tokenizer | None = None, + chat_template: str | None = None, + chat_template_kwargs: Mapping[str, object] | None = None, + ) -> TokenizedTrajectory: ... + + @overload + def tokenize( + self, + *, + multi_history: Literal[True], + model: str | None = None, + base_model: str | None = None, + tokenizer: Tokenizer | None = None, + chat_template: str | None = None, + chat_template_kwargs: Mapping[str, object] | None = None, + ) -> TokenizedMultiHistoryTrajectory: ... + + def tokenize( + self, + *, + multi_history: bool = False, + model: str | None = None, + base_model: str | None = None, + tokenizer: Tokenizer | None = None, + chat_template: str | None = None, + chat_template_kwargs: Mapping[str, object] | None = None, + ) -> TokenizedTrajectory | TokenizedMultiHistoryTrajectory: + from ._tokenize import tokenize_trajectory + + return tokenize_trajectory( + self, + multi_history=multi_history, + model=model, + base_model=base_model, + tokenizer=tokenizer, + chat_template=chat_template, + chat_template_kwargs=chat_template_kwargs, + ) + def messages(self) -> Messages: from ._history import trajectory_messages @@ -527,47 +756,123 @@ def __iter__(self) -> Iterator[Trajectory]: # ty: ignore[invalid-method-overrid def __len__(self) -> int: return len(self.trajectories) + @overload + def tokenize( + self, + *, + multi_history: Literal[False] = False, + model: str | None = None, + base_model: str | None = None, + tokenizer: Tokenizer | None = None, + chat_template: str | None = None, + chat_template_kwargs: Mapping[str, object] | None = None, + ) -> TokenizedTrajectoryGroup[TokenizedTrajectory]: ... -class TokenizedTrajectory(pydantic.BaseModel): + @overload + def tokenize( + self, + *, + multi_history: Literal[True], + model: str | None = None, + base_model: str | None = None, + tokenizer: Tokenizer | None = None, + chat_template: str | None = None, + chat_template_kwargs: Mapping[str, object] | None = None, + ) -> TokenizedTrajectoryGroup[TokenizedMultiHistoryTrajectory]: ... + + def tokenize( + self, + *, + multi_history: bool = False, + model: str | None = None, + base_model: str | None = None, + tokenizer: Tokenizer | None = None, + chat_template: str | None = None, + chat_template_kwargs: Mapping[str, object] | None = None, + ) -> ( + TokenizedTrajectoryGroup[TokenizedTrajectory] + | TokenizedTrajectoryGroup[TokenizedMultiHistoryTrajectory] + ): + from ._tokenize import tokenize_group + + return tokenize_group( + self, + multi_history=multi_history, + model=model, + base_model=base_model, + tokenizer=tokenizer, + chat_template=chat_template, + chat_template_kwargs=chat_template_kwargs, + ) + + +class TokenizedHistory(pydantic.BaseModel): + model_config = pydantic.ConfigDict(ser_json_inf_nan="strings") + + model: str token_ids: list[int] logprobs: list[float] flags: list[TokenFlag] - underlying: Trajectory + @pydantic.model_validator(mode="after") + def validate_tokenwise_lengths(self) -> TokenizedHistory: + if not (len(self.token_ids) == len(self.logprobs) == len(self.flags)): + raise ValueError("Tokenized history fields differ in length") + return self -class TokenizedTrajectoryGroup(pydantic.BaseModel): - trajectories: list[TokenizedTrajectory] - underlying: TrajectoryGroup + +class TokenizedTrajectory(TokenizedHistory): + reward: float + metrics: dict[str, float | int | bool] + metadata: dict[str, MetadataValue] + + +class TokenizedMultiHistoryTrajectory(pydantic.BaseModel): + histories: list[TokenizedHistory] + reward: float + metrics: dict[str, float | int | bool] + metadata: dict[str, MetadataValue] + + +TokenizedTrajectoryT = TypeVar( + "TokenizedTrajectoryT", TokenizedTrajectory, TokenizedMultiHistoryTrajectory +) + + +class TokenizedTrajectoryGroup(pydantic.BaseModel, Generic[TokenizedTrajectoryT]): + trajectories: list[TokenizedTrajectoryT] + metrics: dict[str, float | int | bool] + metadata: dict[str, MetadataValue] @overload -def current_trajectory(*, required: Literal[True]) -> Trajectory: ... +def current_trajectory(*, require: Literal[True]) -> Trajectory: ... @overload -def current_trajectory(*, required: Literal[False] = False) -> Trajectory | None: ... +def current_trajectory(*, require: Literal[False] = False) -> Trajectory | None: ... -def current_trajectory(*, required: bool = False) -> Trajectory | None: +def current_trajectory(*, require: bool = False) -> Trajectory | None: from ._scope import get_current_trajectory - return get_current_trajectory(required=required) + return get_current_trajectory(required=require) @overload -def current_trajectory_group(*, required: Literal[True]) -> TrajectoryGroup: ... +def current_trajectory_group(*, require: Literal[True]) -> TrajectoryGroup: ... @overload def current_trajectory_group( - *, required: Literal[False] = False + *, require: Literal[False] = False ) -> TrajectoryGroup | None: ... -def current_trajectory_group(*, required: bool = False) -> TrajectoryGroup | None: +def current_trajectory_group(*, require: bool = False) -> TrajectoryGroup | None: from ._scope import get_current_trajectory_group - return get_current_trajectory_group(required=required) + return get_current_trajectory_group(required=require) async def trajectory(coroutine: Coroutine[Any, Any, object]) -> Trajectory: @@ -589,98 +894,19 @@ async def trajectory_group( ) -def tokenize_trajectory( - trajectory: Trajectory, - *, - base_model: str | None = None, - model: str | None = None, - chat_template: str | None = None, - chat_template_kwargs: Mapping[str, object] | None = None, -) -> TokenizedTrajectory: - from ._tokenize import tokenize_one - - return tokenize_one( - trajectory, - base_model, - model=model, - chat_template=chat_template, - chat_template_kwargs=chat_template_kwargs, - ) - - -def tokenize_trajectories( - trajectories: Iterable[Trajectory], - *, - base_model: str | None = None, - model: str | None = None, - chat_template: str | None = None, - chat_template_kwargs: Mapping[str, object] | None = None, -) -> list[TokenizedTrajectory]: - return [ - tokenize_trajectory( - item, - base_model=base_model, - model=model, - chat_template=chat_template, - chat_template_kwargs=chat_template_kwargs, - ) - for item in trajectories - ] - - -def tokenize_trajectory_group( - group: TrajectoryGroup, - *, - base_model: str | None = None, - model: str | None = None, - chat_template: str | None = None, - chat_template_kwargs: Mapping[str, object] | None = None, -) -> TokenizedTrajectoryGroup: - return TokenizedTrajectoryGroup( - trajectories=tokenize_trajectories( - group, - base_model=base_model, - model=model, - chat_template=chat_template, - chat_template_kwargs=chat_template_kwargs, - ), - underlying=group, - ) - - -def tokenize_trajectory_groups( - groups: Iterable[TrajectoryGroup], - *, - base_model: str | None = None, - model: str | None = None, - chat_template: str | None = None, - chat_template_kwargs: Mapping[str, object] | None = None, -) -> list[TokenizedTrajectoryGroup]: - return [ - tokenize_trajectory_group( - group, - base_model=base_model, - model=model, - chat_template=chat_template, - chat_template_kwargs=chat_template_kwargs, - ) - for group in groups - ] - - @overload @deprecated("Use current_trajectory() instead.") -def auto_trajectory(*, required: Literal[True]) -> Trajectory: ... +def auto_trajectory(*, require: Literal[True]) -> Trajectory: ... @overload @deprecated("Use current_trajectory() instead.") -def auto_trajectory(*, required: Literal[False] = False) -> Trajectory | None: ... +def auto_trajectory(*, require: Literal[False] = False) -> Trajectory | None: ... @deprecated("Use current_trajectory() instead.") -def auto_trajectory(*, required: bool = False) -> Trajectory | None: - return current_trajectory(required=required) +def auto_trajectory(*, require: bool = False) -> Trajectory | None: + return current_trajectory(require=require) @deprecated("Use trajectory() instead.") @@ -708,25 +934,32 @@ def get_messages(messages_and_choices: MessagesAndChoices) -> Messages: "TrajectoryExchanges", "PydanticException", "History", + "LegacyHistory", + "ChatCompletionsMessageSource", + "AnthropicMessageSource", + "ResponsesItemSource", + "CompletionsSource", + "CompletionsTokenSourceSpan", + "CompletionsStringSourceSpan", "ChatCompletionsHistory", "AnthropicMessagesHistory", "ResponsesHistory", - "CompletionsHistory", + "CompletionsTokenHistory", + "CompletionsStringHistory", "TrajectoryHistory", "Trajectory", "TrajectoryGroup", "TokenizedTrajectory", + "TokenizedHistory", + "TokenizedMultiHistoryTrajectory", "TokenizedTrajectoryGroup", + "Tokenizer", "TokenFlag", "MetadataValue", "current_trajectory", "current_trajectory_group", "trajectory", "trajectory_group", - "tokenize_trajectory", - "tokenize_trajectories", - "tokenize_trajectory_group", - "tokenize_trajectory_groups", "auto_trajectory", "capture_auto_trajectory", "get_messages", diff --git a/src/art/trajectories/_capture/core.py b/src/art/trajectories/_capture/core.py index 8de4f154b..3e4cac320 100644 --- a/src/art/trajectories/_capture/core.py +++ b/src/art/trajectories/_capture/core.py @@ -22,6 +22,43 @@ _adapter_active: contextvars.ContextVar[bool] = contextvars.ContextVar( "art_capture_adapter_active", default=False ) +_SSE_DELIMITERS = (b"\r\n\r\n", b"\n\n", b"\r\r") + + +def _terminal_sse_event(endpoint: Endpoint, block: bytes) -> bool: + try: + lines = block.decode("utf-8").splitlines() + except UnicodeDecodeError: + return False + event_name: str | None = None + data_lines: list[str] = [] + for line in lines: + field, separator, value = line.partition(":") + if not separator: + continue + if value.startswith(" "): + value = value[1:] + if field == "event": + event_name = value + elif field == "data": + data_lines.append(value) + data = "\n".join(data_lines) + if endpoint in {"chat_completions", "completions"}: + return data == "[DONE]" + if endpoint == "responses" and event_name == "response.completed": + return True + if endpoint == "messages" and event_name == "message_stop": + return True + try: + payload = json.loads(data) + except json.JSONDecodeError: + return False + if not isinstance(payload, dict): + return False + event_type = payload.get("type") + return (endpoint == "responses" and event_type == "response.completed") or ( + endpoint == "messages" and event_type == "message_stop" + ) @dataclass @@ -33,10 +70,35 @@ class CaptureState: status_code: int | None = None body: bytearray = field(default_factory=bytearray) captured: bool = False + _event_start: int = field(default=0, init=False, repr=False) + _scan_start: int = field(default=0, init=False, repr=False) def add(self, chunk: bytes) -> None: if not self.captured: self.body.extend(chunk) + if self.request.get("stream") is True and self._reached_terminal_event(): + self.finish() + + def _reached_terminal_event(self) -> bool: + while True: + boundaries = [ + (index, len(delimiter)) + for delimiter in _SSE_DELIMITERS + if (index := self.body.find(delimiter, self._scan_start)) >= 0 + ] + if not boundaries: + self._scan_start = max(self._event_start, len(self.body) - 3) + return False + index, delimiter_length = min(boundaries) + block = bytes(self.body[self._event_start : index]) + self._event_start = index + delimiter_length + self._scan_start = self._event_start + if _terminal_sse_event(self.endpoint, block): + return True + + def discard(self) -> None: + self.body.clear() + self.captured = True def finish(self) -> None: if self.captured: diff --git a/src/art/trajectories/_capture/httpx.py b/src/art/trajectories/_capture/httpx.py index 31293405c..01c2f7efc 100644 --- a/src/art/trajectories/_capture/httpx.py +++ b/src/art/trajectories/_capture/httpx.py @@ -4,12 +4,14 @@ import httpx from httpx._client import UseClientDefault +from httpx._decoders import ContentDecoder from httpx._types import AuthTypes from typing_extensions import TypedDict, Unpack from .core import CaptureState, begin, reset _STATE = "_art_trajectory_capture" +_RAW_ACTIVE = "_art_trajectory_capture_raw_active" class _SendOptions(TypedDict, total=False): @@ -18,13 +20,37 @@ class _SendOptions(TypedDict, total=False): follow_redirects: bool | UseClientDefault +def _capture_decoder(response: httpx.Response) -> ContentDecoder: + shadow = httpx.Response( + response.status_code, + headers=response.headers, + stream=httpx.ByteStream(b""), + ) + return shadow._get_content_decoder() + + +def _capture_preloaded( + state: CaptureState, response: httpx.Response, *, stream: bool +) -> None: + if stream and not response.is_stream_consumed: + return + try: + content = response.content + except httpx.ResponseNotRead: + if not stream: + raise + return + state.add(content) + state.finish() + + def install() -> None: if getattr(httpx.Client.send, "_art_capture", False): return original_send = httpx.Client.send original_async_send = httpx.AsyncClient.send - original_iter = httpx.Response.iter_bytes - original_aiter = httpx.Response.aiter_bytes + original_iter = httpx.Response.iter_raw + original_aiter = httpx.Response.aiter_raw original_close = httpx.Response.close original_aclose = httpx.Response.aclose @@ -45,9 +71,7 @@ def send( if state is not None: state.status_code = response.status_code setattr(response, _STATE, state) - if not kwargs.get("stream", False): - state.add(response.content) - state.finish() + _capture_preloaded(state, response, stream=kwargs.get("stream", False)) return response async def async_send( @@ -67,58 +91,110 @@ async def async_send( if state is not None: state.status_code = response.status_code setattr(response, _STATE, state) - if not kwargs.get("stream", False): - state.add(response.content) - state.finish() + _capture_preloaded(state, response, stream=kwargs.get("stream", False)) return response - def iter_bytes( + def iter_raw( self: httpx.Response, chunk_size: int | None = None ) -> Iterator[bytes]: state: CaptureState | None = getattr(self, _STATE, None) + if state is None: + yield from original_iter(self, chunk_size) + return + try: + decoder = _capture_decoder(self) + except Exception: + state.discard() + yield from original_iter(self, chunk_size) + return completed = False + usable = True + setattr(self, _RAW_ACTIVE, True) try: for chunk in original_iter(self, chunk_size): - if state is not None: - state.add(chunk) + if usable: + try: + state.add(decoder.decode(chunk)) + except Exception: + state.discard() + usable = False yield chunk completed = True finally: - if state is not None and (completed or state.request.get("stream") is True): + if usable and (completed or state.request.get("stream") is True): + try: + state.add(decoder.flush()) + except Exception: + state.discard() + usable = False + setattr(self, _RAW_ACTIVE, False) + if usable and (completed or state.request.get("stream") is True): state.finish() - async def aiter_bytes( + async def aiter_raw( self: httpx.Response, chunk_size: int | None = None ) -> AsyncIterator[bytes]: state: CaptureState | None = getattr(self, _STATE, None) + if state is None: + async for chunk in original_aiter(self, chunk_size): + yield chunk + return + try: + decoder = _capture_decoder(self) + except Exception: + state.discard() + async for chunk in original_aiter(self, chunk_size): + yield chunk + return completed = False + usable = True + setattr(self, _RAW_ACTIVE, True) try: async for chunk in original_aiter(self, chunk_size): - if state is not None: - state.add(chunk) + if usable: + try: + state.add(decoder.decode(chunk)) + except Exception: + state.discard() + usable = False yield chunk completed = True finally: - if state is not None and (completed or state.request.get("stream") is True): + if usable and (completed or state.request.get("stream") is True): + try: + state.add(decoder.flush()) + except Exception: + state.discard() + usable = False + setattr(self, _RAW_ACTIVE, False) + if usable and (completed or state.request.get("stream") is True): state.finish() def close(self: httpx.Response) -> None: original_close(self) state: CaptureState | None = getattr(self, _STATE, None) - if state is not None and state.request.get("stream") is True: + if ( + state is not None + and not getattr(self, _RAW_ACTIVE, False) + and state.request.get("stream") is True + ): state.finish() async def aclose(self: httpx.Response) -> None: await original_aclose(self) state: CaptureState | None = getattr(self, _STATE, None) - if state is not None and state.request.get("stream") is True: + if ( + state is not None + and not getattr(self, _RAW_ACTIVE, False) + and state.request.get("stream") is True + ): state.finish() setattr(send, "_art_capture", True) setattr(async_send, "_art_capture", True) setattr(httpx.Client, "send", send) setattr(httpx.AsyncClient, "send", async_send) - setattr(httpx.Response, "iter_bytes", iter_bytes) - setattr(httpx.Response, "aiter_bytes", aiter_bytes) + setattr(httpx.Response, "iter_raw", iter_raw) + setattr(httpx.Response, "aiter_raw", aiter_raw) setattr(httpx.Response, "close", close) setattr(httpx.Response, "aclose", aclose) diff --git a/src/art/trajectories/_capture/requests.py b/src/art/trajectories/_capture/requests.py index 9af7bc94c..8ff248cb3 100644 --- a/src/art/trajectories/_capture/requests.py +++ b/src/art/trajectories/_capture/requests.py @@ -45,8 +45,8 @@ def iter_content( ): if state is not None: if isinstance(chunk, str): - chunk = chunk.encode(self.encoding or "utf-8") - if isinstance(chunk, bytes): + state.add(chunk.encode(self.encoding or "utf-8")) + elif isinstance(chunk, bytes): state.add(chunk) yield chunk completed = True diff --git a/src/art/trajectories/_history.py b/src/art/trajectories/_history.py index 48300348b..41b246d89 100644 --- a/src/art/trajectories/_history.py +++ b/src/art/trajectories/_history.py @@ -1,23 +1,49 @@ from __future__ import annotations -from collections.abc import Iterable, Mapping, Sequence +from collections.abc import Callable, Iterable, Mapping, Sequence import copy +from dataclasses import dataclass, replace from datetime import datetime -from typing import Protocol, TypeVar +from fnmatch import fnmatchcase +import json +from typing import Generic, Protocol, TypeVar, cast +from anthropic.types import ( + MessageParam as AnthropicMessageParam, +) +from anthropic.types import ( + TextBlockParam as AnthropicTextBlockParam, +) +from anthropic.types import ( + ToolUnionParam as AnthropicToolParam, +) +from openai.types import CompletionChoice +from openai.types.responses import ResponseInputParam +from openai.types.responses import ToolParam as ResponsesToolParam +from openai.types.responses.response_create_params import ( + Conversation as ResponsesConversation, +) +from openai.types.responses.response_input_param import ResponseInputItemParam import pydantic from ..types import Message, Messages, Tools from . import ( AnthropicMessagesHistory, + AnthropicMessageSource, ChatCompletionsExchange, ChatCompletionsHistory, + ChatCompletionsMessageSource, CompletionsExchange, - CompletionsHistory, - History, + CompletionsSource, + CompletionsStringHistory, + CompletionsStringSourceSpan, + CompletionsTokenHistory, + CompletionsTokenSourceSpan, + LegacyHistory, MessagesExchange, ResponsesExchange, ResponsesHistory, + ResponsesItemSource, Trajectory, TrajectoryHistory, ) @@ -31,73 +57,326 @@ class _ModelledExchange(Protocol): def model(self) -> str | None: ... +class _Indexed(Protocol): + index: int + + _ExchangeT = TypeVar("_ExchangeT", bound=_ModelledExchange) +_IndexedT = TypeVar("_IndexedT", bound=_Indexed) _ItemT = TypeVar("_ItemT") +_SourceT = TypeVar("_SourceT") +_ContextT = TypeVar("_ContextT") _MESSAGES = pydantic.TypeAdapter(Messages) _MESSAGE = pydantic.TypeAdapter(Message) _TOOLS = pydantic.TypeAdapter(Tools | None) +_CHAT_KWARGS = pydantic.TypeAdapter(dict[str, object] | None) +_ANTHROPIC_SYSTEM = pydantic.TypeAdapter(str | list[AnthropicTextBlockParam] | None) +_ANTHROPIC_TOOLS = pydantic.TypeAdapter(list[AnthropicToolParam] | None) +_RESPONSE_TOOLS = pydantic.TypeAdapter(list[ResponsesToolParam] | None) +_RESPONSE_CONVERSATION = pydantic.TypeAdapter(ResponsesConversation | None) + + +@dataclass(frozen=True) +class _ChatContext: + tools: Tools | None + template: str | None + kwargs: dict[str, object] | None + + +@dataclass(frozen=True) +class _AnthropicContext: + system: str | list[AnthropicTextBlockParam] | None + tools: list[AnthropicToolParam] | None + template: str | None + kwargs: dict[str, object] | None + + +@dataclass(frozen=True) +class _ResponsesContext: + instructions: str | None + tools: list[ResponsesToolParam] | None + conversation: ResponsesConversation | None + previous_response_id: str | None + template: str | None + kwargs: dict[str, object] | None + + +type _ResponsesRecord = tuple[ + ResponseInputParam, + list[ResponsesItemSource | None], + ResponsesConversation | None, + str | None, +] + + +@dataclass(frozen=True) +class _ResponsesGeneration: + prompt_token_ids: list[int] + output_token_ids: list[int] + output_indices: list[int] -def _select( +@dataclass +class _Branch(Generic[_ItemT, _SourceT, _ContextT]): + items: list[_ItemT] + sources: list[_SourceT | None] + context: _ContextT + order: tuple[int, ...] + first_time: datetime + context_source: _ModelledExchange | None + lineage_keys: list[object] | None = None + + +def _selected_models( exchanges: Sequence[_ExchangeT], model: str | None, protocol: str -) -> list[_ExchangeT]: +) -> list[tuple[str, list[_ExchangeT]]]: + has_exact_match = model is not None and any( + exchange.model == model for exchange in exchanges + ) selected = [ - exchange for exchange in exchanges if model is None or exchange.model == model + exchange + for exchange in exchanges + if model is None + or ( + exchange.model == model + if has_exact_match + else _model_matches(exchange.model, model) + ) ] if not selected: suffix = f" for model {model!r}" if model is not None else "" raise ValueError(f"Trajectory contains no {protocol} exchanges{suffix}") - models = {exchange.model for exchange in selected} - if None in models: + if any(exchange.model is None for exchange in selected): raise ValueError(f"Every {protocol} exchange must identify its model") - if len(models) != 1: + grouped: dict[str, list[_ExchangeT]] = {} + for exchange in sorted(selected, key=lambda item: (item.start_time, item.end_time)): + if exchange.model is None: + raise AssertionError("model identity was checked above") + grouped.setdefault(exchange.model, []).append(exchange) + return list(grouped.items()) + + +def _model_matches(candidate: str | None, pattern: str) -> bool: + return candidate is not None and ( + candidate == pattern or fnmatchcase(candidate, pattern) + ) + + +def _one(history: Sequence[_ItemT], protocol: str) -> _ItemT: + if len(history) != 1: raise ValueError( - f"{protocol} history requires exactly one model; pass model= to select one" + f"{protocol} requires exactly one history; found {len(history)}" ) - return sorted( - selected, key=lambda exchange: (exchange.start_time, exchange.end_time) - ) + return history[0] -def _require_constant(values: Iterable[object], field: str) -> None: - iterator = iter(values) - first = next(iterator) - if any(value != first for value in iterator): - raise ValueError(f"Exchanges with different {field} form different histories") +def _ordered_choices(choices: Sequence[_IndexedT], *, protocol: str) -> list[_IndexedT]: + if not choices: + raise ValueError(f"{protocol} response contains no choices") + indices = [choice.index for choice in choices] + if any( + not isinstance(index, int) or isinstance(index, bool) or index < 0 + for index in indices + ): + raise ValueError(f"{protocol} choice indices must be non-negative integers") + if len(set(indices)) != len(indices): + raise ValueError(f"{protocol} response contains duplicate choice indices") + return sorted(choices, key=lambda choice: choice.index) -def _require_context( - requests: Sequence[Mapping[str, object]], fields: Sequence[str] -) -> None: - for field in fields: - _require_constant((request.get(field) for request in requests), field) +def _is_prefix(prefix: Sequence[object], value: Sequence[object]) -> bool: + return len(prefix) <= len(value) and all( + left == right for left, right in zip(prefix, value[: len(prefix)], strict=True) + ) + + +def _lineage_prompt_sources( + branches: Sequence[_Branch[_ItemT, _SourceT, _ContextT]], + *, + prompt: Sequence[_ItemT], + defaults: Sequence[_SourceT | None], + equivalent: Callable[[_ItemT, _ItemT], bool], + prompt_keys: Sequence[object] | None = None, +) -> list[_SourceT | None]: + candidates = [ + branch + for branch in branches + if len(branch.items) < len(prompt) + and ( + list(prompt_keys[: len(branch.items)]) == branch.lineage_keys + if prompt_keys is not None and branch.lineage_keys is not None + else all( + equivalent(existing, current) + for existing, current in zip(branch.items, prompt, strict=False) + ) + ) + ] + if not candidates: + return list(defaults) + parent = min(candidates, key=lambda branch: (-len(branch.items), branch.order)) + return [*parent.sources, *defaults[len(parent.items) :]] -def _extend( - history: list[_ItemT] | None, +def _extend_branches( + branches: list[_Branch[_ItemT, _SourceT, _ContextT]], + *, prompt: Sequence[_ItemT], - completion: Sequence[_ItemT], - protocol: str, -) -> list[_ItemT]: - if history is None: - history = copy.deepcopy(list(prompt)) - elif len(prompt) < len(history) or list(prompt[: len(history)]) != history: - raise ValueError(f"{protocol} exchanges do not form one append-only history") + prompt_sources: Sequence[_SourceT | None], + outputs: Sequence[tuple[int, Sequence[_ItemT], Sequence[_SourceT | None]]], + context: _ContextT, + sequence: int, + start_time: datetime, + context_source: _ModelledExchange | None = None, + continuation: Callable[[_Branch[_ItemT, _SourceT, _ContextT]], bool] | None = None, + prompt_lineage_keys: Sequence[object] | None = None, + output_lineage_keys: Sequence[Sequence[object]] | None = None, +) -> None: + if len(prompt) != len(prompt_sources): + raise AssertionError("prompt sources must parallel prompt items") + if (prompt_lineage_keys is None) != (output_lineage_keys is None): + raise AssertionError("prompt and output lineage keys must be provided together") + if prompt_lineage_keys is not None and len(prompt_lineage_keys) != len(prompt): + raise AssertionError("prompt lineage keys must parallel prompt items") + if output_lineage_keys is not None and len(output_lineage_keys) != len(outputs): + raise AssertionError("output lineage keys must parallel outputs") + if output_lineage_keys is not None and any( + len(keys) != len(output) + for keys, (_, output, _) in zip(output_lineage_keys, outputs, strict=True) + ): + raise AssertionError("output lineage keys must parallel output items") + candidates = [ + (index, branch) + for index, branch in enumerate(branches) + if branch.context == context + and _is_prefix(branch.items, prompt) + and (continuation is None or continuation(branch)) + ] + parent: _Branch[_ItemT, _SourceT, _ContextT] | None = None + remove_parent = False + if candidates: + index, parent = min( + candidates, + key=lambda item: (-len(item[1].items), item[1].order), + ) + remove_parent = True + branches.pop(index) + sources = [ + *parent.sources, + *prompt_sources[len(parent.items) :], + ] else: - history.extend(copy.deepcopy(list(prompt[len(history) :]))) - history.extend(copy.deepcopy(list(completion))) - return history + regenerations = [ + branch + for branch in branches + if branch.context == context + and _is_prefix(prompt, branch.items) + and (continuation is None or continuation(branch)) + ] + if regenerations: + parent = min(regenerations, key=lambda branch: branch.order) + sources = copy.copy(parent.sources[: len(prompt)]) + else: + sources = copy.copy(prompt_sources) + + base_order = parent.order if parent is not None else (sequence,) + first_time = parent.first_time if parent is not None else start_time + created = [ + _Branch( + items=( + [*prompt, *output] + if prompt_lineage_keys is not None + else [*copy.deepcopy(prompt), *copy.deepcopy(output)] + ), + sources=[*sources, *output_sources], + context=copy.deepcopy(context), + order=(*base_order, choice_index), + first_time=first_time, + context_source=context_source, + lineage_keys=( + [*prompt_lineage_keys, *output_lineage_keys[position]] + if prompt_lineage_keys is not None and output_lineage_keys is not None + else None + ), + ) + for position, (choice_index, output, output_sources) in enumerate(outputs) + ] + if not created and remove_parent and parent is not None: + branches.append(parent) + branches.extend(created) -def _only_choice(exchange: ChatCompletionsExchange | CompletionsExchange) -> None: - if len(exchange.response.choices) != 1: - raise ValueError("Multiple response choices form multiple histories") +def _contains_tokens(tokens: Sequence[int], sampled: Sequence[int]) -> bool: + return any( + list(tokens[start : start + len(sampled)]) == list(sampled) + for start in range(len(tokens) - len(sampled) + 1) + ) + + +def _chat_retains_sampled_reasoning( + branch: _Branch[Message, ChatCompletionsMessageSource, _ChatContext], + exchange: ChatCompletionsExchange, + prompt_length: int, +) -> bool: + """Split histories when a template drops a prior sampled reasoning span.""" + + if not exchange.response.choices: + return True + from ._tokenize import _chat_choice_tokens + + response_data = exchange.response.model_dump(mode="python") + prompt_ids, _, _ = _chat_choice_tokens(exchange.response.choices[0], response_data) + if prompt_ids is None: + return True + for message, source in zip( + branch.items[:prompt_length], branch.sources[:prompt_length], strict=True + ): + reasoning = message.get("reasoning") or message.get("reasoning_content") + if ( + not reasoning + or source is None + or source.choice_index is None + or not isinstance(source.exchange, ChatCompletionsExchange) + ): + continue + source_response = source.exchange.response + choice = next( + ( + item + for item in source_response.choices + if item.index == source.choice_index + ), + None, + ) + if choice is None: + continue + _, sampled_ids, _ = _chat_choice_tokens( + choice, source_response.model_dump(mode="python") + ) + if sampled_ids and not _contains_tokens(prompt_ids, sampled_ids): + return False + return True -def _model(exchange: _ModelledExchange) -> str: - if exchange.model is None: - raise AssertionError("_select returned an exchange without a model") - return exchange.model +def _chat_message_key(message: Message, *, visible_only: bool = False) -> str: + data = normalize_chat_message(message) + if visible_only: + data.pop("reasoning", None) + data.pop("reasoning_content", None) + return json.dumps(data, sort_keys=True, default=str) + + +def normalize_chat_message(message: Mapping[str, object]) -> dict[str, object]: + """Return the canonical history form of an OpenAI chat message.""" + + data = copy.deepcopy(dict(message)) + data.pop("annotations", None) + if data.get("role") == "assistant" and data.get("content") is None: + data["content"] = "" + elif data.get("content") is None: + data.pop("content", None) + if data.get("tool_calls") == []: + data.pop("tool_calls") + return data def _require_unmixed(trajectory: Trajectory) -> None: @@ -111,201 +390,995 @@ def _require_unmixed(trajectory: Trajectory) -> None: ) -def legacy_as_chat_completions_history(history: History) -> ChatCompletionsHistory: +def legacy_as_chat_completions_history( + history: LegacyHistory, *, model: str | None +) -> ChatCompletionsHistory: + messages = history.messages() return ChatCompletionsHistory( - model=None, - messages=history.messages(), + model=model, + messages=messages, + message_sources=[None] * len(messages), tools=copy.deepcopy(history.tools), ) -def chat_completions_history( +def chat_completions_histories( trajectory: Trajectory, *, model: str | None -) -> ChatCompletionsHistory: +) -> list[ChatCompletionsHistory]: _require_unmixed(trajectory) if not trajectory.exchanges: - if trajectory.additional_histories: - raise ValueError("Trajectory contains multiple legacy histories") - return ChatCompletionsHistory( - model=model, - messages=History( - messages_and_choices=trajectory.messages_and_choices, - tools=trajectory.tools, - ).messages(), - tools=copy.deepcopy(trajectory.tools), - ) - exchanges = _select( + return [ + history.as_chat_completions_history(model=model) + for history in [ + LegacyHistory( + messages_and_choices=trajectory.messages_and_choices, + tools=trajectory.tools, + ), + *trajectory.additional_histories, + ] + ] + + histories: list[ChatCompletionsHistory] = [] + for selected_model, exchanges in _selected_models( trajectory.exchanges.chat_completions, model, "Chat Completions" - ) - _require_context( - [exchange.request for exchange in exchanges], - ("tools", "chat_template", "chat_template_kwargs", "cache_salt"), - ) - messages: Messages | None = None - for exchange in exchanges: - _only_choice(exchange) - prompt = _MESSAGES.validate_python(exchange.request.get("messages", [])) - response = _MESSAGE.validate_python( - exchange.response.choices[0].message.model_dump( - mode="python", exclude_none=True + ): + branches: list[ + _Branch[Message, ChatCompletionsMessageSource, _ChatContext] + ] = [] + for sequence, exchange in enumerate(exchanges): + prompt = [ + cast(Message, normalize_chat_message(message)) + for message in _MESSAGES.validate_python( + exchange.request.get("messages", []) + ) + ] + prompt_lineage_keys = [ + _chat_message_key(message, visible_only=True) for message in prompt + ] + prompt_sources = _lineage_prompt_sources( + branches, + prompt=prompt, + defaults=[ + ChatCompletionsMessageSource(exchange=exchange, request_index=index) + for index in range(len(prompt)) + ], + equivalent=lambda left, right: ( + _chat_message_key(left, visible_only=True) + == _chat_message_key(right, visible_only=True) + ), + prompt_keys=prompt_lineage_keys, + ) + outputs: list[ + tuple[ + int, + list[Message], + list[ChatCompletionsMessageSource | None], + ] + ] = [] + for choice in _ordered_choices( + exchange.response.choices, protocol="Chat Completions" + ): + response = _MESSAGE.validate_python( + normalize_chat_message( + choice.message.model_dump(mode="python", exclude_none=True) + ) + ) + response_source = ChatCompletionsMessageSource( + exchange=exchange, choice_index=choice.index + ) + outputs.append( + ( + choice.index, + [response], + [response_source], + ) + ) + output_lineage_keys = [ + [_chat_message_key(output[0], visible_only=True)] + for _, output, _ in outputs + ] + _extend_branches( + branches, + prompt=prompt, + prompt_sources=prompt_sources, + outputs=outputs, + context=_ChatContext( + tools=_TOOLS.validate_python(exchange.request.get("tools")), + template=exchange.request.get("chat_template"), + kwargs=_CHAT_KWARGS.validate_python( + exchange.request.get("chat_template_kwargs") + ), + ), + sequence=sequence, + start_time=exchange.start_time, + continuation=lambda branch: _chat_retains_sampled_reasoning( + branch, exchange, len(prompt) + ), + prompt_lineage_keys=prompt_lineage_keys, + output_lineage_keys=output_lineage_keys, + ) + for branch in sorted(branches, key=lambda item: (item.first_time, item.order)): + histories.append( + ChatCompletionsHistory( + model=selected_model, + messages=copy.deepcopy(branch.items), + message_sources=copy.copy(branch.sources), + tools=copy.deepcopy(branch.context.tools), + chat_template=branch.context.template, + chat_template_kwargs=copy.deepcopy(branch.context.kwargs), + ) ) + return histories + + +def chat_completions_history( + trajectory: Trajectory, *, model: str | None +) -> ChatCompletionsHistory: + histories = chat_completions_histories(trajectory, model=model) + if model is None and len({history.model for history in histories}) > 1: + raise ValueError( + "Chat Completions history requires exactly one model; pass model= to select one" ) - messages = _extend(messages, prompt, [response], "Chat Completions") - first = exchanges[0].request - return ChatCompletionsHistory( - model=_model(exchanges[0]), - messages=messages or [], - tools=copy.deepcopy(first.get("tools")), - chat_template=first.get("chat_template"), - chat_template_kwargs=copy.deepcopy(first.get("chat_template_kwargs")), - ) + return _one(histories, "Chat Completions history") + + +def anthropic_messages_histories( + trajectory: Trajectory, *, model: str | None +) -> list[AnthropicMessagesHistory]: + _require_unmixed(trajectory) + histories: list[AnthropicMessagesHistory] = [] + for selected_model, exchanges in _selected_models( + trajectory.exchanges.messages, model, "Anthropic Messages" + ): + branches: list[ + _Branch[AnthropicMessageParam, AnthropicMessageSource, _AnthropicContext] + ] = [] + for sequence, exchange in enumerate(exchanges): + prompt = _anthropic_prompt(exchange.request.get("messages", [])) + prompt_sources = _lineage_prompt_sources( + branches, + prompt=prompt, + defaults=[ + AnthropicMessageSource(exchange=exchange, request_index=index) + for index in range(len(prompt)) + ], + equivalent=lambda left, right: ( + _anthropic_message_key(left, visible_only=True) + == _anthropic_message_key(right, visible_only=True) + ), + ) + response = cast( + AnthropicMessageParam, + { + "role": "assistant", + "content": [ + block.model_dump(mode="json", exclude_none=True) + for block in exchange.response.content + ], + }, + ) + response_source = AnthropicMessageSource(exchange=exchange) + _extend_branches( + branches, + prompt=prompt, + prompt_sources=prompt_sources, + outputs=[ + ( + 0, + [response], + [response_source], + ) + ], + context=_AnthropicContext( + system=_ANTHROPIC_SYSTEM.validate_python( + exchange.request.get("system") + ), + tools=_ANTHROPIC_TOOLS.validate_python( + exchange.request.get("tools") + ), + template=exchange.request.get("chat_template"), + kwargs=_CHAT_KWARGS.validate_python( + exchange.request.get("chat_template_kwargs") + ), + ), + sequence=sequence, + start_time=exchange.start_time, + context_source=( + exchange if exchange.request.get("system") is not None else None + ), + ) + for branch in sorted(branches, key=lambda item: (item.first_time, item.order)): + histories.append( + AnthropicMessagesHistory( + model=selected_model, + messages=copy.deepcopy(branch.items), + message_sources=copy.copy(branch.sources), + system=copy.deepcopy(branch.context.system), + system_source=( + cast(MessagesExchange, branch.context_source) + if branch.context_source is not None + else None + ), + tools=copy.deepcopy(branch.context.tools), + chat_template=branch.context.template, + chat_template_kwargs=copy.deepcopy(branch.context.kwargs), + ) + ) + return histories + + +def _anthropic_prompt(value: object) -> list[AnthropicMessageParam]: + if not isinstance(value, list): + raise ValueError("Anthropic messages must be a list") + messages: list[AnthropicMessageParam] = [] + for message in value: + if not isinstance(message, Mapping): + raise ValueError("Anthropic messages must be JSON objects") + content = message.get("content") + if not isinstance(content, (str, list)): + raise ValueError("Anthropic message content must be text or a list") + if isinstance(content, list) and any( + not isinstance(block, (Mapping, pydantic.BaseModel)) for block in content + ): + raise ValueError("Anthropic message content blocks must be JSON objects") + messages.append(cast(AnthropicMessageParam, copy.deepcopy(dict(message)))) + return messages + + +def _anthropic_message_key( + message: AnthropicMessageParam, *, visible_only: bool = False +) -> str: + normalized = copy.deepcopy(dict(message)) + content = message.get("content") + if isinstance(content, str): + blocks: list[dict[str, object]] = [{"type": "text", "text": content}] + elif isinstance(content, list): + blocks = [] + for block in content: + if isinstance(block, pydantic.BaseModel): + data = block.model_dump(mode="json", exclude_none=True) + elif isinstance(block, Mapping): + data = copy.deepcopy(dict(block)) + else: + raise ValueError( + "Anthropic message content blocks must be JSON objects" + ) + kind = data.get("type") + if visible_only and kind in {"thinking", "redacted_thinking"}: + continue + for field in ("token_ids", "logprobs"): + data.pop(field, None) + blocks.append(data) + else: + raise ValueError("Anthropic message content must be text or a list") + normalized["content"] = blocks + return json.dumps(normalized, sort_keys=True, default=str) def anthropic_messages_history( trajectory: Trajectory, *, model: str | None ) -> AnthropicMessagesHistory: - _require_unmixed(trajectory) - exchanges = _select(trajectory.exchanges.messages, model, "Anthropic Messages") - _require_context( - [exchange.request for exchange in exchanges], - ( - "system", - "tools", - "thinking", - "chat_template", - "chat_template_kwargs", - "cache_salt", - ), - ) - messages: list[object] | None = None - for exchange in exchanges: - prompt = copy.deepcopy(exchange.request.get("messages", [])) - response = { - "role": "assistant", - "content": [ - block.model_dump(mode="python", exclude_none=True) - for block in exchange.response.content - ], - } - messages = _extend(messages, prompt, [response], "Anthropic Messages") - first = exchanges[0].request - return AnthropicMessagesHistory( - model=_model(exchanges[0]), - system=copy.deepcopy(first.get("system")), - messages=messages or [], - tools=copy.deepcopy(first.get("tools")), - thinking=copy.deepcopy(first.get("thinking")), - chat_template=first.get("chat_template"), - chat_template_kwargs=copy.deepcopy(first.get("chat_template_kwargs")), - ) + histories = anthropic_messages_histories(trajectory, model=model) + if model is None and len({history.model for history in histories}) > 1: + raise ValueError( + "Anthropic Messages history requires exactly one model; pass model= to select one" + ) + return _one(histories, "Anthropic Messages history") -def _responses_input(value: object) -> list[object]: +def _responses_input(value: object) -> ResponseInputParam: if isinstance(value, str): - return [{"role": "user", "content": value}] + return [cast(ResponseInputItemParam, {"role": "user", "content": value})] if isinstance(value, list): - return [copy.deepcopy(item) for item in value] + return [_copy_response_item(item) for item in value] if value is None: return [] raise ValueError("Responses input must be text or a list of input items") -def responses_history(trajectory: Trajectory, *, model: str | None) -> ResponsesHistory: +def _copy_response_item(value: object) -> ResponseInputItemParam: + if not isinstance(value, dict): + raise ValueError("Responses input items must be JSON objects") + # OpenAI models this as a large TypedDict union. Capture already validated + # the wire shape; this boundary makes the detached protocol-native copy. + return cast(ResponseInputItemParam, copy.deepcopy(value)) + + +def responses_histories( + trajectory: Trajectory, *, model: str | None +) -> list[ResponsesHistory]: _require_unmixed(trajectory) - exchanges = _select(trajectory.exchanges.responses, model, "Responses") - _require_context( - [exchange.request for exchange in exchanges], + histories: list[ResponsesHistory] = [] + for selected_model, exchanges in _selected_models( + trajectory.exchanges.responses, model, "Responses" + ): + branches: list[ + _Branch[ResponseInputItemParam, ResponsesItemSource, _ResponsesContext] + ] = [] + responses: dict[str, _ResponsesRecord] = {} + conversations: dict[str, _ResponsesRecord] = {} + for sequence, exchange in enumerate(exchanges): + request = exchange.request + prompt = _responses_input(request.get("input")) + prompt_sources: list[ResponsesItemSource | None] = [ + ResponsesItemSource(exchange=exchange, request_index=index) + for index in range(len(prompt)) + ] + requested_conversation = _RESPONSE_CONVERSATION.validate_python( + request.get("conversation") + ) + previous = request.get("previous_response_id") + inherited_conversation: ResponsesConversation | None = None + inherited_previous: str | None = None + prior: _ResponsesRecord | None = None + if isinstance(previous, str) and previous in responses: + prior = responses[previous] + elif requested_conversation is not None: + prior = conversations.get( + json.dumps(requested_conversation, sort_keys=True, default=str) + ) + if prior is not None: + ( + prior_items, + prior_sources, + inherited_conversation, + inherited_previous, + ) = prior + prompt = [*copy.deepcopy(prior_items), *prompt] + prompt_sources = [*copy.copy(prior_sources), *prompt_sources] + prompt_sources = _lineage_prompt_sources( + branches, + prompt=prompt, + defaults=prompt_sources, + equivalent=lambda left, right: ( + _response_item_key(left) == _response_item_key(right) + ), + ) + output = [ + cast( + ResponseInputItemParam, + item.model_dump(mode="json", exclude_none=True), + ) + for item in exchange.response.output + ] + generations, generation_indices = _responses_generations( + exchange, output=output + ) + output_sources: list[ResponsesItemSource | None] = [ + ResponsesItemSource( + exchange=exchange, + output_index=index, + generation_index=generation_indices[index], + ) + for index in range(len(output)) + ] + instructions = request.get("instructions") + if instructions is not None and not isinstance(instructions, str): + raise ValueError("Responses instructions must be text") + external_previous = ( + previous + if isinstance(previous, str) and previous not in responses + else inherited_previous + ) + conversation = ( + requested_conversation + if requested_conversation is not None + else inherited_conversation + ) + context = _ResponsesContext( + instructions=instructions, + tools=_RESPONSE_TOOLS.validate_python(request.get("tools")), + conversation=conversation, + previous_response_id=external_previous, + template=request.get("chat_template"), + kwargs=_CHAT_KWARGS.validate_python( + request.get("chat_template_kwargs") + ), + ) + final_items = [*copy.deepcopy(prompt), *copy.deepcopy(output)] + final_sources = [*copy.copy(prompt_sources), *copy.copy(output_sources)] + if generations: + for generation_index, generation in enumerate(generations): + outputless = not generation.output_indices + if outputless and generation_index != len(generations) - 1: + raise ValueError( + "A nonterminal Responses token generation without " + "native output items cannot be projected as a history" + ) + if outputless: + output_start = output_end = len(output) + else: + output_start = generation.output_indices[0] + output_end = generation.output_indices[-1] + 1 + if not outputless and generation_index == len(generations) - 1: + output_end = len(output) + generation_prompt = [ + *copy.deepcopy(prompt), + *copy.deepcopy(output[:output_start]), + ] + generation_prompt_sources = [ + *copy.copy(prompt_sources), + *copy.copy(output_sources[:output_start]), + ] + continuation = lambda branch, generation=generation: ( + _responses_generation_extends( + branch, + generation=generation, + ) + ) + extends = any( + branch.context == context + and _is_prefix(branch.items, generation_prompt) + and continuation(branch) + for branch in branches + ) + if not extends: + retained = [ + (item, source) + for item, source in zip( + generation_prompt, + generation_prompt_sources, + strict=True, + ) + if not ( + item.get("type") == "reasoning" + and source is not None + and source.output_index is not None + and source.generation_index is not None + ) + ] + generation_prompt = [item for item, _ in retained] + generation_prompt_sources = [ + ( + replace(source, generation_index=None) + if source is not None + and source.exchange is exchange + and source.generation_index is not None + else source + ) + for _, source in retained + ] + if outputless: + generation_output = [ + cast( + ResponseInputItemParam, + {"role": "assistant", "content": ""}, + ) + ] + generation_output_sources = [ + ResponsesItemSource( + exchange=exchange, + generation_index=generation_index, + ) + ] + else: + generation_output = output[output_start:output_end] + generation_output_sources = output_sources[ + output_start:output_end + ] + _extend_branches( + branches, + prompt=generation_prompt, + prompt_sources=generation_prompt_sources, + outputs=[ + ( + generation_index, + generation_output, + generation_output_sources, + ) + ], + context=context, + sequence=sequence, + start_time=exchange.start_time, + context_source=(exchange if instructions is not None else None), + continuation=continuation, + ) + final_items = [ + *copy.deepcopy(generation_prompt), + *copy.deepcopy(generation_output), + ] + final_sources = [ + *copy.copy(generation_prompt_sources), + *copy.copy(generation_output_sources), + ] + else: + _extend_branches( + branches, + prompt=prompt, + prompt_sources=prompt_sources, + outputs=[(0, output, output_sources)], + context=context, + sequence=sequence, + start_time=exchange.start_time, + context_source=exchange if instructions is not None else None, + ) + record = ( + final_items, + final_sources, + copy.deepcopy(conversation), + external_previous, + ) + responses[exchange.response.id] = record + if conversation is not None: + conversations[json.dumps(conversation, sort_keys=True, default=str)] = ( + record + ) + for branch in sorted(branches, key=lambda item: (item.first_time, item.order)): + histories.append( + ResponsesHistory( + model=selected_model, + input=copy.deepcopy(branch.items), + input_sources=copy.copy(branch.sources), + instructions=branch.context.instructions, + instructions_source=( + cast(ResponsesExchange, branch.context_source) + if branch.context_source is not None + else None + ), + tools=copy.deepcopy(branch.context.tools), + conversation=copy.deepcopy(branch.context.conversation), + previous_response_id=branch.context.previous_response_id, + chat_template=branch.context.template, + chat_template_kwargs=copy.deepcopy(branch.context.kwargs), + ) + ) + return histories + + +def _response_item_key(item: ResponseInputItemParam) -> str: + return json.dumps(item, sort_keys=True, default=str) + + +def _responses_generations( + exchange: ResponsesExchange, *, output: Sequence[ResponseInputItemParam] +) -> tuple[list[_ResponsesGeneration], list[int | None]]: + from ._tokenize import _response_generations + + generations = _response_generations(exchange.response) + result: list[int | None] = [None] * len(output) + parsed: list[_ResponsesGeneration] = [] + for generation_index, generation in enumerate(generations): + if generation.prompt_token_ids is None: + raise ValueError( + "Responses generation without an exact prompt cannot be projected" + ) + if generation.output_token_ids is None: + raise ValueError( + "Responses generation without exact output tokens cannot yet be " + "projected as a history" + ) + for output_index in generation.output_indices: + result[output_index] = generation_index + parsed.append( + _ResponsesGeneration( + prompt_token_ids=generation.prompt_token_ids, + output_token_ids=generation.output_token_ids, + output_indices=generation.output_indices, + ) + ) + return parsed, result + + +def _responses_generation_extends( + branch: _Branch[ResponseInputItemParam, ResponsesItemSource, _ResponsesContext], + *, + generation: _ResponsesGeneration, +) -> bool: + prior_source = next( ( - "instructions", - "tools", - "chat_template", - "chat_template_kwargs", - "cache_salt", + source + for source in reversed(branch.sources) + if source is not None and source.generation_index is not None ), + None, ) - items: list[object] | None = None - previous_response_id: str | None = None - for exchange in exchanges: - request = exchange.request - prompt = _responses_input(request.get("input")) - previous = request.get("previous_response_id") - if previous is not None: - if previous != previous_response_id or items is None: - raise ValueError( - "Responses exchange refers to a response outside this history" - ) - prompt = [*items, *prompt] - output = [ - item.model_dump(mode="python", exclude_none=True) - for item in exchange.response.output - ] - items = _extend(items, prompt, output, "Responses") - previous_response_id = exchange.response.id - first = exchanges[0].request - return ResponsesHistory( - model=_model(exchanges[0]), - input=items or [], - instructions=first.get("instructions"), - tools=copy.deepcopy(first.get("tools")), - chat_template=first.get("chat_template"), - chat_template_kwargs=copy.deepcopy(first.get("chat_template_kwargs")), + if prior_source is None or prior_source.generation_index is None: + return True + prior_exchange = prior_source.exchange + prior_output = [ + cast( + ResponseInputItemParam, + item.model_dump(mode="json", exclude_none=True), + ) + for item in prior_exchange.response.output + ] + prior_generations, _ = _responses_generations(prior_exchange, output=prior_output) + if not 0 <= prior_source.generation_index < len(prior_generations): + raise ValueError("Responses generation source index is out of bounds") + prior = prior_generations[prior_source.generation_index] + if _is_prefix( + [*prior.prompt_token_ids, *prior.output_token_ids], + generation.prompt_token_ids, + ): + return True + return not any( + getattr(prior_exchange.response.output[index], "type", None) == "reasoning" + for index in prior.output_indices ) -def completions_history( - trajectory: Trajectory, *, model: str | None -) -> CompletionsHistory: +def responses_history(trajectory: Trajectory, *, model: str | None) -> ResponsesHistory: + histories = responses_histories(trajectory, model=model) + if model is None and len({history.model for history in histories}) > 1: + raise ValueError( + "Responses history requires exactly one model; pass model= to select one" + ) + return _one(histories, "Responses history") + + +def _completion_exact_tokens( + exchange: CompletionsExchange, choice_index: int +) -> tuple[list[int] | None, list[int] | None]: from ._tokenize import _completion_tokens, _exact_token_ids - _require_unmixed(trajectory) - exchanges = _select(trajectory.exchanges.completions, model, "Completions") - _require_context([exchange.request for exchange in exchanges], ("cache_salt",)) - token_ids: list[int] = [] - sampled_spans: list[tuple[int, int]] = [] - for index, exchange in enumerate(exchanges): - _only_choice(exchange) - if exchange.request.get("echo") is True: - raise ValueError("Completions history does not support echo=True") - prompt, completion, _ = _completion_tokens(exchange.response) - request_prompt = exchange.request.get("prompt") - if prompt is None and isinstance(request_prompt, list): + choice = next( + choice for choice in exchange.response.choices if choice.index == choice_index + ) + prompt, completion, _ = _completion_tokens( + exchange.response.model_copy(update={"choices": [choice]}), + echo=exchange.request.get("echo") is True, + ) + request_prompt = exchange.request.get("prompt") + if prompt is None and isinstance(request_prompt, list): + try: prompt = _exact_token_ids( request_prompt, field="Completions request prompt" ) - if prompt is None or completion is None: + except ValueError: + pass + return prompt, completion + + +def completions_token_histories( + trajectory: Trajectory, *, model: str | None +) -> list[CompletionsTokenHistory]: + _require_unmixed(trajectory) + histories: list[CompletionsTokenHistory] = [] + for selected_model, exchanges in _selected_models( + trajectory.exchanges.completions, model, "Completions" + ): + branches: list[_Branch[int, CompletionsSource, tuple[()]]] = [] + complete = True + for sequence, exchange in enumerate(exchanges): + if exchange.request.get("suffix") is not None: + raise ValueError("Completions suffix is not supported") + choice_groups = _completion_choice_groups(exchange) + for prompt_index, request_prompt in enumerate( + _completion_prompts(exchange.request.get("prompt")) + ): + choices = choice_groups[prompt_index] + for choice in choices: + prompt, completion = _completion_exact_tokens( + exchange, choice.index + ) + if prompt is None: + if not isinstance(request_prompt, list): + complete = False + continue + prompt = copy.deepcopy(request_prompt) + if completion is None: + complete = False + continue + prompt_source = CompletionsSource( + exchange=exchange, prompt_index=prompt_index + ) + output_source = CompletionsSource( + exchange=exchange, + prompt_index=prompt_index, + choice_index=choice.index, + ) + _extend_branches( + branches, + prompt=prompt, + prompt_sources=[prompt_source] * len(prompt), + outputs=[ + ( + choice.index, + completion, + [output_source] * len(completion), + ) + ], + context=(), + sequence=sequence, + start_time=exchange.start_time, + ) + if not complete: + raise ValueError( + "Completions token history requires exact token IDs for every choice" + ) + for branch in branches: + histories.append( + CompletionsTokenHistory( + model=selected_model, + prompt=copy.deepcopy(branch.items), + prompt_sources=_token_source_spans(branch.sources), + sampled_spans=_sampled_spans(branch.sources), + ) + ) + if not histories: + raise ValueError("Completions token history requires exact token IDs") + return histories + + +def completions_token_history( + trajectory: Trajectory, *, model: str | None +) -> CompletionsTokenHistory: + return _one( + completions_token_histories(trajectory, model=model), + "Completions token history", + ) + + +def _completion_prompts(value: object) -> list[str | list[int]]: + if isinstance(value, str): + return [value] + if isinstance(value, list): + if all(isinstance(item, int) and not isinstance(item, bool) for item in value): + return [_checked_token_ids(value)] + prompts: list[str | list[int]] = [] + for item in value: + if isinstance(item, str): + prompts.append(item) + elif isinstance(item, list) and all( + isinstance(token, int) and not isinstance(token, bool) for token in item + ): + prompts.append(_checked_token_ids(item)) + else: + raise ValueError("Invalid batched Completions prompt") + return prompts + raise ValueError("Completions prompt must be text or token IDs") + + +def _checked_token_ids(value: Sequence[object]) -> list[int]: + result: list[int] = [] + for token in value: + if not isinstance(token, int) or isinstance(token, bool) or token < 0: + raise ValueError( + "Completions token prompts must contain non-negative integers" + ) + result.append(token) + return result + + +def _completion_choice_groups( + exchange: CompletionsExchange, +) -> list[list[CompletionChoice]]: + prompts = _completion_prompts(exchange.request.get("prompt")) + count = len(prompts) + choices = _ordered_choices(exchange.response.choices, protocol="Completions") + missing_prompt_index = object() + explicit_matches: dict[int, int] = {} + for choice in choices: + raw_prompt_index = (choice.model_extra or {}).get( + "prompt_index", missing_prompt_index + ) + if raw_prompt_index is missing_prompt_index: + continue + if ( + not isinstance(raw_prompt_index, int) + or isinstance(raw_prompt_index, bool) + or not 0 <= raw_prompt_index < count + ): raise ValueError( - "Completions history requires exact prompt and output token IDs" + "Completions choice prompt_index must identify a batched prompt" ) - if index == 0: - token_ids.extend(prompt) - elif len(prompt) < len(token_ids) or prompt[: len(token_ids)] != token_ids: + explicit_matches[choice.index] = raw_prompt_index + if count == 1: + return [choices] + requested_n = exchange.request.get("n") + if requested_n is not None and ( + not isinstance(requested_n, int) + or isinstance(requested_n, bool) + or requested_n < 1 + ): + raise ValueError("Completions n must be a positive integer") + + exact_matches: dict[int, int] = {} + for choice in choices: + prompt_ids, _ = _completion_exact_tokens(exchange, choice.index) + if prompt_ids is None: + continue + matches = [ + prompt_index + for prompt_index, prompt in enumerate(prompts) + if isinstance(prompt, list) and prompt == prompt_ids + ] + explicit = explicit_matches.get(choice.index) + if ( + explicit is not None + and isinstance(prompts[explicit], list) + and prompts[explicit] != prompt_ids + ): raise ValueError( - "Completions exchanges do not form one append-only token history" + "Completions prompt_index contradicts exact prompt evidence" ) - else: - token_ids.extend(prompt[len(token_ids) :]) - start = len(token_ids) - token_ids.extend(completion) - sampled_spans.append((start, len(token_ids))) - return CompletionsHistory( - model=_model(exchanges[0]), - token_ids=token_ids, - sampled_spans=sampled_spans, + if len(matches) == 1: + exact_matches[choice.index] = matches[0] + + evidence_matches = {**exact_matches, **explicit_matches} + if len(evidence_matches) == len(choices): + groups = [[] for _ in prompts] + for choice in choices: + groups[evidence_matches[choice.index]].append(choice) + if all(groups) and ( + requested_n is None or all(len(group) == requested_n for group in groups) + ): + return groups + raise ValueError("Cannot associate Completions choices with batched prompts") + + if len(choices) % count: + raise ValueError("Cannot associate Completions choices with batched prompts") + per_prompt = len(choices) // count + if requested_n is not None and requested_n != per_prompt: + raise ValueError("Cannot associate Completions choices with batched prompts") + if [choice.index for choice in choices] != list(range(len(choices))): + raise ValueError("Ambiguous Completions choice-to-prompt association") + groups = [ + choices[prompt_index * per_prompt : (prompt_index + 1) * per_prompt] + for prompt_index in range(count) + ] + for prompt_index, group in enumerate(groups): + if any( + choice.index in evidence_matches + and evidence_matches[choice.index] != prompt_index + for choice in group + ): + raise ValueError("Completions prompt evidence contradicts choice indices") + return groups + + +def completions_string_histories( + trajectory: Trajectory, *, model: str | None +) -> list[CompletionsStringHistory]: + _require_unmixed(trajectory) + histories: list[CompletionsStringHistory] = [] + for selected_model, exchanges in _selected_models( + trajectory.exchanges.completions, model, "Completions" + ): + branches: list[_Branch[str, CompletionsSource, tuple[()]]] = [] + complete = True + for sequence, exchange in enumerate(exchanges): + if exchange.request.get("suffix") is not None: + raise ValueError("Completions suffix is not supported") + choice_groups = _completion_choice_groups(exchange) + for prompt_index, prompt in enumerate( + _completion_prompts(exchange.request.get("prompt")) + ): + if not isinstance(prompt, str): + complete = False + continue + for choice in choice_groups[prompt_index]: + text = choice.text + if exchange.request.get("echo") is True: + if not text.startswith(prompt): + raise ValueError( + "Cannot locate echoed Completions prompt boundary" + ) + text = text[len(prompt) :] + prompt_source = CompletionsSource( + exchange=exchange, prompt_index=prompt_index + ) + output_source = CompletionsSource( + exchange=exchange, + prompt_index=prompt_index, + choice_index=choice.index, + ) + _extend_branches( + branches, + prompt=list(prompt), + prompt_sources=[prompt_source] * len(prompt), + outputs=[ + ( + choice.index, + list(text), + [output_source] * len(text), + ) + ], + context=(), + sequence=sequence, + start_time=exchange.start_time, + ) + if not complete: + raise ValueError( + "Completions string history requires text prompts for every choice" + ) + histories.extend( + CompletionsStringHistory( + model=selected_model, + prompt="".join(branch.items), + prompt_sources=_string_source_spans(branch.sources), + sampled_spans=_sampled_spans(branch.sources), + ) + for branch in branches + ) + if not histories: + raise ValueError("Completions string history requires text prompts") + return histories + + +def completions_string_history( + trajectory: Trajectory, *, model: str | None +) -> CompletionsStringHistory: + return _one( + completions_string_histories(trajectory, model=model), + "Completions string history", ) +def _source_runs( + sources: Sequence[CompletionsSource | None], +) -> Iterable[tuple[int, int, CompletionsSource | None]]: + if not sources: + return + start = 0 + source = sources[0] + for index, item in enumerate(sources[1:], 1): + if item != source: + yield start, index, source + start, source = index, item + yield start, len(sources), source + + +def _token_source_spans( + sources: Sequence[CompletionsSource | None], +) -> list[CompletionsTokenSourceSpan]: + return [ + CompletionsTokenSourceSpan(start=start, end=end, source=source) + for start, end, source in _source_runs(sources) + ] + + +def _string_source_spans( + sources: Sequence[CompletionsSource | None], +) -> list[CompletionsStringSourceSpan]: + return [ + CompletionsStringSourceSpan(start=start, end=end, source=source) + for start, end, source in _source_runs(sources) + ] + + +def _sampled_spans( + sources: Sequence[CompletionsSource | None], +) -> list[tuple[int, int]]: + return [ + (start, end) + for start, end, source in _source_runs(sources) + if source is not None and source.choice_index is not None + ] + + def anthropic_as_chat_completions_history( history: AnthropicMessagesHistory, ) -> ChatCompletionsHistory: from ._tokenize import _anthropic_messages, _openai_tools - messages = _anthropic_messages( - {"system": history.system, "messages": history.messages} - ) + messages: list[dict[str, object]] = [] + sources: list[ChatCompletionsMessageSource | None] = [] + if history.system: + messages.extend(_anthropic_messages({"system": history.system, "messages": []})) + sources.append( + ChatCompletionsMessageSource(exchange=history.system_source) + if history.system_source is not None + else None + ) + for message, source in zip(history.messages, history.message_sources, strict=True): + converted_messages = _anthropic_messages({"messages": [message]}) + messages.extend(converted_messages) + converted = ( + ChatCompletionsMessageSource( + exchange=source.exchange, + request_index=source.request_index, + output_indices=(0,) if source.request_index is None else None, + ) + if source is not None + else None + ) + sources.extend([converted] * len(converted_messages)) tools = _TOOLS.validate_python(_openai_tools(history.tools, dialect="messages")) return ChatCompletionsHistory( model=history.model, messages=_MESSAGES.validate_python(messages), + message_sources=sources, tools=tools, chat_template=history.chat_template, chat_template_kwargs=copy.deepcopy(history.chat_template_kwargs), @@ -315,65 +1388,268 @@ def anthropic_as_chat_completions_history( def responses_as_chat_completions_history( history: ResponsesHistory, ) -> ChatCompletionsHistory: + if history.previous_response_id is not None or history.conversation is not None: + raise ValueError( + "Opaque Responses context cannot be represented as Chat Completions" + ) from ._tokenize import _openai_tools, _responses_messages - messages = _responses_messages( - {"instructions": history.instructions, "input": history.input} - ) + messages = _responses_messages({"instructions": history.instructions, "input": []}) + sources: list[ChatCompletionsMessageSource | None] = [] + if history.instructions: + sources.append( + ChatCompletionsMessageSource(exchange=history.instructions_source) + if history.instructions_source is not None + else None + ) + + def converted( + contributors: Sequence[ResponsesItemSource | None], + ) -> ChatCompletionsMessageSource | None: + if not contributors or any(source is None for source in contributors): + return None + present = [source for source in contributors if source is not None] + exchanges = {id(source.exchange): source.exchange for source in present} + if len(exchanges) != 1: + raise ValueError( + "One projected Chat message cannot span multiple Responses exchanges" + ) + generation_indices = { + source.generation_index + for source in present + if source.generation_index is not None + } + if len(generation_indices) > 1: + raise ValueError( + "One projected Chat message cannot span multiple Responses generations" + ) + output_indices = tuple( + dict.fromkeys( + source.output_index + for source in present + if source.output_index is not None + ) + ) + generation_index = next(iter(generation_indices), None) + first = present[0] + if output_indices or generation_index is not None: + return ChatCompletionsMessageSource( + exchange=first.exchange, + output_indices=output_indices, + generation_index=generation_index, + ) + return ChatCompletionsMessageSource( + exchange=first.exchange, + request_index=next( + source.request_index + for source in present + if source.request_index is not None + ), + ) + + no_output = object() + + def compatible( + contributors: Sequence[ResponsesItemSource | None], + source: ResponsesItemSource | None, + ) -> bool: + def sampled(item: ResponsesItemSource | None) -> bool: + return item is not None and ( + item.output_index is not None or item.generation_index is not None + ) + + present = [item for item in contributors if item is not None] + if ( + source is not None + and present + and source.exchange is not present[0].exchange + ): + return False + if contributors and sampled(source) != any( + sampled(item) for item in contributors + ): + return False + if source is None: + return True + + def output_generation(item: ResponsesItemSource) -> object: + if item.output_index is None and item.generation_index is None: + return no_output + return item.generation_index + + generations = { + value + for item in present + if (value := output_generation(item)) is not no_output + } + candidate = output_generation(source) + return candidate is no_output or not generations or candidate in generations + + item_groups: list[list[ResponseInputItemParam]] = [] + source_groups: list[list[ResponsesItemSource | None]] = [] + pending_reasoning_items: list[ResponseInputItemParam] = [] + pending_reasoning: list[ResponsesItemSource | None] = [] + tool_message_items: list[ResponseInputItemParam] | None = None + tool_message_sources: list[ResponsesItemSource | None] | None = None + for item, source in zip(history.input, history.input_sources, strict=True): + kind = item.get("type") + if kind == "reasoning": + if pending_reasoning and not compatible(pending_reasoning, source): + item_groups.append([*pending_reasoning_items]) + source_groups.append([*pending_reasoning]) + pending_reasoning_items.clear() + pending_reasoning.clear() + tool_message_items = None + tool_message_sources = None + pending_reasoning_items.append(item) + pending_reasoning.append(source) + continue + if kind == "function_call": + if tool_message_sources is None: + if pending_reasoning and not compatible(pending_reasoning, source): + item_groups.append([*pending_reasoning_items]) + source_groups.append([*pending_reasoning]) + pending_reasoning_items.clear() + pending_reasoning.clear() + tool_message_items = [*pending_reasoning_items, item] + tool_message_sources = [*pending_reasoning, source] + item_groups.append(tool_message_items) + source_groups.append(tool_message_sources) + pending_reasoning_items.clear() + pending_reasoning.clear() + elif compatible(tool_message_sources, source): + assert tool_message_items is not None + tool_message_items.append(item) + tool_message_sources.append(source) + else: + tool_message_items = [item] + tool_message_sources = [source] + item_groups.append(tool_message_items) + source_groups.append(tool_message_sources) + continue + tool_message_items = None + tool_message_sources = None + if pending_reasoning: + if ( + kind in {None, "message"} + and item.get("role") == "assistant" + and compatible(pending_reasoning, source) + ): + item_groups.append([*pending_reasoning_items, item]) + source_groups.append([*pending_reasoning, source]) + else: + item_groups.append([*pending_reasoning_items]) + source_groups.append([*pending_reasoning]) + item_groups.append([item]) + source_groups.append([source]) + pending_reasoning_items.clear() + pending_reasoning.clear() + else: + item_groups.append([item]) + source_groups.append([source]) + if pending_reasoning: + item_groups.append(pending_reasoning_items) + source_groups.append(pending_reasoning) + for items, group in zip(item_groups, source_groups, strict=True): + projected = _responses_messages({"input": items}) + if len(projected) != 1: + raise AssertionError( + "One Responses projection group must produce one Chat message" + ) + messages.extend(projected) + sources.append(converted(group)) + if len(sources) != len(messages): + raise AssertionError("Responses conversion sources must parallel messages") tools = _TOOLS.validate_python(_openai_tools(history.tools, dialect="responses")) return ChatCompletionsHistory( model=history.model, messages=_MESSAGES.validate_python(messages), + message_sources=sources, tools=tools, chat_template=history.chat_template, chat_template_kwargs=copy.deepcopy(history.chat_template_kwargs), ) -def trajectory_history( +def trajectory_histories( trajectory: Trajectory, *, model: str | None -) -> TrajectoryHistory: +) -> list[TrajectoryHistory]: _require_unmixed(trajectory) if not trajectory.exchanges: - if trajectory.additional_histories: - raise ValueError("Trajectory contains multiple legacy histories") - if model is not None: - raise ValueError("Legacy trajectory histories do not identify a model") - return History( - messages_and_choices=copy.deepcopy(trajectory.messages_and_choices), - tools=copy.deepcopy(trajectory.tools), - ) + return [ + LegacyHistory( + messages_and_choices=copy.deepcopy(trajectory.messages_and_choices), + tools=copy.deepcopy(trajectory.tools), + ), + *copy.deepcopy(trajectory.additional_histories), + ] - candidates = [ - name - for name, exchanges in ( - ("chat_completions", trajectory.exchanges.chat_completions), - ("completions", trajectory.exchanges.completions), - ("responses", trajectory.exchanges.responses), - ("messages", trajectory.exchanges.messages), - ) - if any(model is None or exchange.model == model for exchange in exchanges) + all_exchanges = [ + *trajectory.exchanges.chat_completions, + *trajectory.exchanges.completions, + *trajectory.exchanges.responses, + *trajectory.exchanges.messages, ] + has_exact_match = model is not None and any( + exchange.model == model for exchange in all_exchanges + ) + + def is_selected(exchange: _ModelledExchange) -> bool: + if model is None: + return True + if has_exact_match: + return exchange.model == model + return _model_matches(exchange.model, model) + + candidates: list[TrajectoryHistory] = [] + if any(is_selected(exchange) for exchange in trajectory.exchanges.chat_completions): + candidates.extend(chat_completions_histories(trajectory, model=model)) + if any(is_selected(exchange) for exchange in trajectory.exchanges.completions): + try: + candidates.extend(completions_token_histories(trajectory, model=model)) + except ValueError as error: + if "requires exact token IDs" not in str(error): + raise + candidates.extend(completions_string_histories(trajectory, model=model)) + if any(is_selected(exchange) for exchange in trajectory.exchanges.responses): + candidates.extend(responses_histories(trajectory, model=model)) + if any(is_selected(exchange) for exchange in trajectory.exchanges.messages): + candidates.extend(anthropic_messages_histories(trajectory, model=model)) if not candidates: suffix = f" for model {model!r}" if model is not None else "" raise ValueError(f"Trajectory contains no exchanges{suffix}") - if len(candidates) != 1: + protocols = {type(history) for history in candidates} + if len(protocols) != 1: raise ValueError( "Trajectory resolves to multiple protocol histories; use a protocol-specific method" ) - protocol = candidates[0] - if protocol == "chat_completions": - return chat_completions_history(trajectory, model=model) - if protocol == "completions": - return completions_history(trajectory, model=model) - if protocol == "responses": - return responses_history(trajectory, model=model) - return anthropic_messages_history(trajectory, model=model) + return candidates + + +def trajectory_history( + trajectory: Trajectory, *, model: str | None +) -> TrajectoryHistory: + histories = trajectory_histories(trajectory, model=model) + if ( + model is None + and len( + { + history.model + for history in histories + if not isinstance(history, LegacyHistory) + } + ) + > 1 + ): + raise ValueError( + "Trajectory history requires exactly one model; pass model= to select one" + ) + return _one(histories, "Trajectory history") def trajectory_messages(trajectory: Trajectory) -> Messages: if not trajectory.exchanges: - return History( + return LegacyHistory( messages_and_choices=trajectory.messages_and_choices, tools=trajectory.tools, ).messages() diff --git a/src/art/trajectories/_protocols.py b/src/art/trajectories/_protocols.py index db89be737..594d5a0e3 100644 --- a/src/art/trajectories/_protocols.py +++ b/src/art/trajectories/_protocols.py @@ -2,6 +2,7 @@ from datetime import datetime import json +import math from typing import Any, Literal from urllib.parse import urlsplit @@ -52,9 +53,12 @@ def endpoint_for_url(url: str) -> Endpoint | None: def _sse_events(body: bytes) -> list[tuple[str | None, SSEPayload]]: - text = body.decode("utf-8").replace("\r\n", "\n") + text = body.decode("utf-8").replace("\r\n", "\n").replace("\r", "\n") events: list[tuple[str | None, SSEPayload]] = [] - for block in text.split("\n\n"): + blocks = text.split("\n\n") + if not text.endswith("\n\n"): + blocks.pop() + for block in blocks: event_name: str | None = None data_lines: list[str] = [] for line in block.splitlines(): @@ -221,9 +225,46 @@ def _messages_response(body: bytes, *, stream: bool) -> Message: complete = False token_ids: list[int] = [] logprobs: list[float] = [] + prompt_token_ids: list[int] = [] + block_token_ids: dict[int, list[int]] = {} + block_logprobs: dict[int, list[float]] = {} + + def parsed_token_ids(value: object, field: str) -> list[int] | None: + if value is None: + return None + if not isinstance(value, list): + raise ValueError(f"{field} must contain integer token IDs") + result: list[int] = [] + for item in value: + if not isinstance(item, int) or isinstance(item, bool) or item < 0: + raise ValueError(f"{field} must contain integer token IDs") + result.append(item) + return result + + def token_logprobs(value: object, field: str) -> list[float] | None: + if value is None: + return None + if not isinstance(value, list): + raise ValueError(f"{field} must contain numeric logprobs") + result: list[float] = [] + for item in value: + if ( + not isinstance(item, (int, float)) + or isinstance(item, bool) + or not math.isfinite(item) + ): + raise ValueError(f"{field} must contain numeric logprobs") + result.append(float(item)) + return result + for event_name, payload in _sse_events(body): if not isinstance(payload, dict): continue + event_type = payload.get("type") or event_name + if event_name == "ping" or event_type == "ping": + continue + if event_name == "error" or event_type == "error": + raise ValueError("Anthropic Messages stream returned an error event") if event_name and "type" not in payload: payload = {**payload, "type": event_name} event = adapter.validate_python(payload) @@ -232,25 +273,80 @@ def _messages_response(body: bytes, *, stream: bool) -> Message: current_snapshot=snapshot, output_format=NOT_GIVEN, ) - if event.type == "message_delta": - event_token_ids = payload.get("token_ids") - event_logprobs = payload.get("logprobs") - if isinstance(event_token_ids, list) and all( - isinstance(value, int) for value in event_token_ids - ): + if event.type == "message_start": + message = payload.get("message") + values = ( + message.get("prompt_token_ids") if isinstance(message, dict) else None + ) + if ( + values := parsed_token_ids(values, "Messages prompt_token_ids") + ) is not None: + prompt_token_ids = values + elif event.type == "message_delta": + if ( + values := parsed_token_ids( + payload.get("prompt_token_ids"), "Messages prompt_token_ids" + ) + ) is not None: + prompt_token_ids = values + if ( + event_token_ids := parsed_token_ids( + payload.get("token_ids"), "Messages token_ids" + ) + ) is not None: token_ids = event_token_ids - if isinstance(event_logprobs, list) and all( - isinstance(value, (int, float)) for value in event_logprobs + if ( + event_logprobs := token_logprobs( + payload.get("logprobs"), "Messages logprobs" + ) + ) is not None: + logprobs = event_logprobs + elif event.type in {"content_block_start", "content_block_delta"}: + index = payload.get("index") + event_token_ids = parsed_token_ids( + payload.get("token_ids"), "Messages content token_ids" + ) + event_logprobs = token_logprobs( + payload.get("logprobs"), "Messages content logprobs" + ) + if ( + isinstance(index, int) + and not isinstance(index, bool) + and event_token_ids ): - logprobs = [float(value) for value in event_logprobs] + block_token_ids.setdefault(index, []).extend(event_token_ids) + if ( + isinstance(index, int) + and not isinstance(index, bool) + and event_logprobs + ): + block_logprobs.setdefault(index, []).extend(event_logprobs) complete = complete or event.type == "message_stop" if snapshot is None or not complete: raise ValueError("Incomplete Messages stream") data = snapshot.model_dump(mode="python") + content = data.get("content") + if isinstance(content, list): + rebuilt_content: list[object] = [] + for index, raw_block in enumerate(content): + if not isinstance(raw_block, dict): + rebuilt_content.append(raw_block) + continue + block: dict[str, Any] = { + str(key): value for key, value in raw_block.items() + } + if values := block_token_ids.get(index): + block["token_ids"] = values + if values := block_logprobs.get(index): + block["logprobs"] = values + rebuilt_content.append(block) + data["content"] = rebuilt_content if token_ids: data["token_ids"] = token_ids if logprobs: data["logprobs"] = logprobs + if prompt_token_ids: + data["prompt_token_ids"] = prompt_token_ids return Message.model_validate(data) diff --git a/src/art/trajectories/_selection.py b/src/art/trajectories/_selection.py new file mode 100644 index 000000000..69ccc9d2b --- /dev/null +++ b/src/art/trajectories/_selection.py @@ -0,0 +1,118 @@ +from __future__ import annotations + +from dataclasses import dataclass +from fnmatch import fnmatchcase +import re +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from . import Trajectory + + +@dataclass(frozen=True, slots=True) +class ModelSelector: + """Private model selection policy used at training boundaries.""" + + value: str + automatic_family: tuple[str, str] | None = None + allow_glob: bool = False + + def __post_init__(self) -> None: + if not self.value: + raise ValueError("A model selector cannot be empty") + + def matches(self, candidate: str) -> bool: + if self.automatic_family is not None: + prefix, separator = self.automatic_family + return bool( + re.fullmatch( + f"{re.escape(prefix)}{re.escape(separator)}[0-9]+", + candidate, + ) + ) + return ( + fnmatchcase(candidate, self.value) + if self.allow_glob + else candidate == self.value + ) + + +def public_model_selector(value: str) -> ModelSelector: + return ModelSelector(value, allow_glob=True) + + +def automatic_training_model_selector(value: str) -> ModelSelector: + if match := re.fullmatch(r"(.*)@([0-9]+)", value): + return ModelSelector(value, (match.group(1), "@")) + if match := re.fullmatch(r"(.*):step([0-9]+)", value): + return ModelSelector(value, (match.group(1), ":step")) + return ModelSelector(value) + + +def resolve_training_model( + trajectory: Trajectory, + selector: ModelSelector | str | None, +) -> str: + """Resolve one concrete, single-protocol captured model for training.""" + + exchanges_by_protocol = { + "Chat Completions": trajectory.exchanges.chat_completions, + "Completions": trajectory.exchanges.completions, + "Responses": trajectory.exchanges.responses, + "Anthropic Messages": trajectory.exchanges.messages, + } + exchanges = [ + (protocol, exchange) + for protocol, protocol_exchanges in exchanges_by_protocol.items() + for exchange in protocol_exchanges + ] + if not exchanges: + raise ValueError("Exchange training requires at least one captured exchange") + if any(exchange.model is None for _, exchange in exchanges): + raise ValueError("Every training exchange must identify its model") + + concrete_models = {exchange.model for _, exchange in exchanges} + if selector is None: + matches = concrete_models + else: + selector = ( + public_model_selector(selector) if isinstance(selector, str) else selector + ) + exact = ( + {candidate for candidate in concrete_models if candidate == selector.value} + if selector.automatic_family is None + else set() + ) + matches = exact or { + candidate + for candidate in concrete_models + if candidate is not None and selector.matches(candidate) + } + if not matches: + value = selector.value if isinstance(selector, ModelSelector) else selector + raise ValueError(f"Trajectory contains no exchanges for model {value!r}") + if len(matches) != 1: + raise ValueError( + "Exchange training requires exactly one concrete model; matched " + f"{sorted(matches)}" + ) + selected_model = next(iter(matches)) + if selected_model is None: + raise AssertionError("model identity was checked above") + protocols = { + protocol for protocol, exchange in exchanges if exchange.model == selected_model + } + if len(protocols) != 1: + raise ValueError( + "Exchange training does not support mixed protocols for one model; found " + f"{sorted(protocols)}" + ) + return selected_model + + +__all__ = [ + "ModelSelector", + "automatic_training_model_selector", + "public_model_selector", + "resolve_training_model", +] diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index a02199872..84c420764 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -1,12 +1,20 @@ from __future__ import annotations -from collections.abc import Mapping +from bisect import bisect_left +import codecs +from collections.abc import Hashable, Mapping, Sequence +from copy import deepcopy from dataclasses import dataclass +from datetime import datetime +from functools import lru_cache +from hashlib import sha256 +import json import math import re -from typing import Any, Protocol, cast +from typing import TYPE_CHECKING, Any, Literal, Protocol, TypeVar, cast +import warnings -from anthropic.types import Message +from anthropic.types import Message, MessageParam, TextBlock from openai.types import Completion from openai.types.chat import ChatCompletion from openai.types.chat.chat_completion import Choice @@ -14,17 +22,37 @@ from pydantic import BaseModel from . import ( + AnthropicMessagesHistory, ChatCompletionsExchange, + ChatCompletionsHistory, CompletionsExchange, + CompletionsSource, + CompletionsStringHistory, + CompletionsStringSourceSpan, + CompletionsTokenHistory, + CompletionsTokenSourceSpan, + History, + LegacyHistory, MessagesExchange, ResponsesExchange, + ResponsesHistory, TokenFlag, + TokenizedHistory, + TokenizedMultiHistoryTrajectory, TokenizedTrajectory, + TokenizedTrajectoryGroup, + Tokenizer, Trajectory, + TrajectoryGroup, ) +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 @dataclass @@ -35,31 +63,287 @@ class _TokenizerConfig: chat_template_kwargs: Mapping[str, object] | None = None -class _Tokenizer(Protocol): - def __call__(self, text: str, *, add_special_tokens: bool = False) -> object: ... - - def apply_chat_template( +class _OffsetTokenizer(Protocol): + def __call__( self, - messages: list[dict[str, Any]], + text: str, *, - tools: object, - tokenize: bool, - add_generation_prompt: bool, - chat_template: str | None = None, - **kwargs: object, + add_special_tokens: bool, + return_offsets_mapping: Literal[True], ) -> object: ... -def _as_tokenizer(tokenizer: object) -> _Tokenizer: +@dataclass(frozen=True) +class _SampledOutput: + text: str | None + token_ids: list[int] + start: int + + +@dataclass(frozen=True) +class _SampledSourceKey: + protocol: Literal["chat_completions", "responses", "messages", "completions"] + response_id: str + start_time: datetime + end_time: datetime + index: int + prompt_index: int | None + evidence_fingerprint: str + + +@dataclass +class _HistoryTokenizationTrace: + source_keys: list[_SampledSourceKey | None] + sources: dict[_SampledSourceKey, object] + + def validate(self, tokenized: TokenizedHistory) -> None: + if len(self.source_keys) != len(tokenized.token_ids): + raise AssertionError( + "Tokenization trace differs in length from tokenized data" + ) + for flag, source_key in zip(tokenized.flags, self.source_keys, strict=True): + if bool(flag & TokenFlag.SAMPLED) != (source_key is not None): + raise AssertionError( + "Tokenization trace must identify every sampled token exactly once" + ) + if source_key is not None and source_key not in self.sources: + raise AssertionError("Tokenization trace source key is unresolved") + + +_SourceKeyT = TypeVar("_SourceKeyT", bound=Hashable) + + +def _first_introduction_mask( + source_keys: Sequence[_SourceKeyT | None], + seen: set[_SourceKeyT], +) -> list[bool]: + keys = {key for key in source_keys if key is not None} + new_keys = keys - seen + seen.update(keys) + return [key in new_keys if key is not None else False for key in source_keys] + + +def _require_causal_predecessor(trainable: Sequence[bool]) -> None: + if trainable and trainable[0]: + raise ValueError("A trainable trajectory cannot start with a sampled token") + + +@dataclass +class _TraceBuilder: + trace: _HistoryTokenizationTrace | None = None + + def set( + self, + tokenized: TokenizedHistory, + source_keys: list[_SampledSourceKey | None], + sources: dict[_SampledSourceKey, object], + ) -> None: + trace = _HistoryTokenizationTrace(source_keys=source_keys, sources=sources) + trace.validate(tokenized) + self.trace = trace + + +def _fingerprint(value: object) -> str: + serialized = json.dumps( + value, + ensure_ascii=False, + separators=(",", ":"), + sort_keys=True, + ) + return sha256(serialized.encode()).hexdigest() + + +def _chat_logprob_fingerprint_evidence(choice: Choice) -> dict[str, object] | None: + if choice.logprobs is None: + return None + + def values(items: Sequence[object] | None) -> list[dict[str, object]]: + result: list[dict[str, object]] = [] + for item in items or []: + data = _dump(item) + result.append( + { + key: data[key] + for key in ("token", "token_id", "logprob", "bytes") + if key in data + } + ) + return result + + return { + "content": values(choice.logprobs.content), + "refusal": values(choice.logprobs.refusal), + } + + +def _sampled_evidence_fingerprint( + exchange: Exchange, + *, + protocol: Literal["chat_completions", "responses", "messages", "completions"], + index: int, +) -> str: + if protocol == "chat_completions": + if not isinstance(exchange, ChatCompletionsExchange): + raise TypeError("Chat source has the wrong exchange type") + choice = next( + choice for choice in exchange.response.choices if choice.index == index + ) + choice_extra = choice.model_extra or {} + evidence = { + "message": choice.message.model_dump( + mode="json", + include={ + "role", + "content", + "refusal", + "reasoning", + "reasoning_content", + "tool_calls", + "function_call", + "audio", + }, + exclude_none=True, + ), + "token_ids": choice_extra.get("token_ids"), + "logprobs": _chat_logprob_fingerprint_evidence(choice), + "finish_reason": choice.finish_reason, + } + elif protocol == "completions": + if not isinstance(exchange, CompletionsExchange): + raise TypeError("Completions source has the wrong exchange type") + choice = next( + choice for choice in exchange.response.choices if choice.index == index + ) + choice_extra = choice.model_extra or {} + logprobs = _dump(choice.logprobs) + evidence = { + "text": choice.text, + "token_ids": choice_extra.get("token_ids"), + "logprobs": { + key: logprobs[key] + for key in ("tokens", "token_logprobs") + if key in logprobs + }, + "finish_reason": choice.finish_reason, + } + elif protocol == "responses": + if not isinstance(exchange, ResponsesExchange): + raise TypeError("Responses source has the wrong exchange type") + generations = (exchange.response.model_extra or {}).get("token_generations") + if isinstance(generations, list) and 0 <= index < len(generations): + generation = _string_dict(generations[index]) or {} + evidence = { + key: generation[key] + for key in ("output_tokens", "output_indices") + if key in generation + } + else: + evidence = [ + output.model_dump(mode="json", exclude_none=True) + for output in exchange.response.output + ] + else: + if not isinstance(exchange, MessagesExchange): + raise TypeError("Messages source has the wrong exchange type") + response_extra = exchange.response.model_extra or {} + evidence = { + "content": [ + block.model_dump(mode="json", exclude_none=True) + for block in exchange.response.content + ], + "token_ids": response_extra.get("token_ids"), + "logprobs": response_extra.get("logprobs"), + "stop_reason": exchange.response.stop_reason, + } + return _fingerprint(evidence) + + +def _source_key( + exchange: Exchange, + *, + protocol: Literal["chat_completions", "responses", "messages", "completions"], + index: int, + prompt_index: int | None = None, +) -> _SampledSourceKey: + return _SampledSourceKey( + protocol=protocol, + response_id=str(getattr(exchange.response, "id", "")), + start_time=exchange.start_time, + end_time=exchange.end_time, + index=index, + prompt_index=prompt_index, + # Internal projection may copy an exchange to isolate one choice. A + # source-specific identity remains stable across those copies without + # hashing a growing request or unrelated choices. + evidence_fingerprint=_sampled_evidence_fingerprint( + exchange, protocol=protocol, index=index + ), + ) + + +def _sampled_source_key(source: object) -> _SampledSourceKey: + exchange = getattr(source, "exchange", None) + if isinstance(exchange, ChatCompletionsExchange): + index = getattr(source, "choice_index", None) + if not isinstance(index, int) or isinstance(index, bool): + raise ValueError("Sampled Chat source has no choice index") + return _source_key(exchange, protocol="chat_completions", index=index) + if isinstance(exchange, ResponsesExchange): + index = getattr(source, "generation_index", None) + if index is None and not _response_generations(exchange.response): + index = 0 + if not isinstance(index, int) or isinstance(index, bool): + raise ValueError("Sampled Responses source has no generation identity") + return _source_key(exchange, protocol="responses", index=index) + if isinstance(exchange, MessagesExchange): + return _source_key(exchange, protocol="messages", index=0) + if isinstance(exchange, CompletionsExchange): + index = getattr(source, "choice_index", None) + prompt_index = getattr(source, "prompt_index", None) + if not isinstance(index, int) or isinstance(index, bool): + raise ValueError("Sampled Completions source has no choice index") + if not isinstance(prompt_index, int) or isinstance(prompt_index, bool): + raise ValueError("Sampled Completions source has no prompt index") + return _source_key( + exchange, + protocol="completions", + index=index, + prompt_index=prompt_index, + ) + raise ValueError("Sampled token source has an unsupported exchange") + + +def _exchange_sampled_source_key(exchange: Exchange) -> _SampledSourceKey: + if isinstance(exchange, ChatCompletionsExchange): + return _source_key( + exchange, + protocol="chat_completions", + index=exchange.response.choices[0].index, + ) + if isinstance(exchange, CompletionsExchange): + return _source_key( + exchange, + protocol="completions", + index=exchange.response.choices[0].index, + prompt_index=0, + ) + if isinstance(exchange, ResponsesExchange): + return _source_key(exchange, protocol="responses", index=0) + if isinstance(exchange, MessagesExchange): + return _source_key(exchange, protocol="messages", index=0) + raise TypeError(f"Unsupported sampled exchange: {type(exchange).__name__}") + + +def _as_tokenizer(tokenizer: object) -> Tokenizer: # Transformers' annotation permits only string-valued message dictionaries, # although its runtime API supports the structured content ART must tokenize. # Exact-token paths may only need decode(); fallback paths exercise these # capabilities directly and report the missing method at that point. - return cast(_Tokenizer, tokenizer) + return cast(Tokenizer, tokenizer) def _string_dict(value: object) -> dict[str, Any] | None: - if not isinstance(value, dict) or not all(isinstance(key, str) for key in value): + if not isinstance(value, Mapping) or not all(isinstance(key, str) for key in value): return None return {key: item for key, item in value.items() if isinstance(key, str)} @@ -84,15 +368,25 @@ def _dump(value: object) -> dict[str, Any]: return _string_dict(value) or {} +def _field(value: object, name: str, default: object = None) -> object: + return ( + value.get(name, default) + if isinstance(value, Mapping) + else getattr(value, name, default) + ) + + def _token_id(value: object) -> int | None: - if isinstance(value, int) and not isinstance(value, bool): + if isinstance(value, int) and not isinstance(value, bool) and value >= 0: return value if isinstance(value, str) and (match := _TOKEN_ID.fullmatch(value)): return int(match.group(1)) return None -def _exact_token_ids(values: object, *, field: str) -> list[int] | None: +def _exact_token_ids( + values: object, *, field: str, empty_is_missing: bool = False +) -> list[int] | None: if values is None: return None if not isinstance(values, list): @@ -103,7 +397,7 @@ def _exact_token_ids(values: object, *, field: str) -> list[int] | None: if token_id is None: raise ValueError(f"{field} contains an invalid exact token ID") token_ids.append(token_id) - return token_ids + return None if empty_is_missing and not token_ids else token_ids def _pair_token_id(data: dict[str, Any], *, required: bool, field: str) -> int | None: @@ -161,35 +455,57 @@ def _logprob_values(values: object) -> list[float]: return result -def _chat_choice_tokens( - choice: Choice, response_data: dict[str, Any] -) -> tuple[list[int] | None, list[int] | None, list[float]]: +def _chat_logprob_entries(choice: Choice) -> list[object]: + if choice.logprobs is None: + return [] + return [ + *(choice.logprobs.content or []), + *(choice.logprobs.refusal or []), + ] + + +def _chat_choice_output_tokens( + choice: Choice, +) -> tuple[list[int] | None, list[float]]: choice_data = _dump(choice) - prompt = choice_data.get("prompt_token_ids") - if prompt is None: - prompt = response_data.get("prompt_token_ids") - prompt_ids = _exact_token_ids(prompt, field="Chat Completions prompt_token_ids") token_ids = _exact_token_ids( choice_data.get("token_ids"), field="Chat Completions token_ids", ) - logprob_values = None - if choice.logprobs is not None: - logprob_values = choice.logprobs.content or choice.logprobs.refusal - values = list(logprob_values or []) + values = _chat_logprob_entries(choice) message = _dump(choice.message) if token_ids == [] and ( values or any(message.get(key) for key in ("content", "refusal", "tool_calls")) ): token_ids = None - pair_ids, logprobs = _pairs(values, field="Chat Completions logprobs") + pair_ids, pair_logprobs = _pairs(values, field="Chat Completions logprobs") + positional_logprobs = _logprob_values(values) if token_ids is not None and pair_ids and token_ids != pair_ids: raise ValueError("Response token IDs disagree with choice logprobs") selected = token_ids if token_ids is not None else pair_ids or None + logprobs = pair_logprobs or positional_logprobs + if selected is not None and values and len(logprobs) != len(selected): + raise ValueError("Chat Completions token IDs and logprobs differ in length") + return selected, logprobs or [math.nan] * len(selected or []) + + +def _chat_choice_tokens( + choice: Choice, response_data: dict[str, Any] +) -> tuple[list[int] | None, list[int] | None, list[float]]: + choice_data = _dump(choice) + prompt = choice_data.get("prompt_token_ids") + if prompt is None: + prompt = response_data.get("prompt_token_ids") + prompt_ids = _exact_token_ids( + prompt, + field="Chat Completions prompt_token_ids", + empty_is_missing=True, + ) + selected, logprobs = _chat_choice_output_tokens(choice) return ( prompt_ids, selected, - logprobs or _logprob_values(values) or [math.nan] * len(selected or []), + logprobs, ) @@ -201,9 +517,12 @@ def _chat_tokens( return _chat_choice_tokens(response.choices[0], _dump(response)) -def _completion_tokens( +def _completion_evidence( response: Completion, -) -> tuple[list[int] | None, list[int] | None, list[float]]: + *, + echo: bool = False, + empty_prompt_is_exact: bool = False, +) -> tuple[list[int] | None, list[int] | None, list[float], list[float]]: if len(response.choices) != 1: raise ValueError("Trajectory tokenization requires exactly one response choice") choice = response.choices[0] @@ -212,7 +531,11 @@ def _completion_tokens( prompt = choice_data.get("prompt_token_ids") if prompt is None: prompt = response_data.get("prompt_token_ids") - prompt_ids = _exact_token_ids(prompt, field="Completions prompt_token_ids") + prompt_ids = _exact_token_ids( + prompt, + field="Completions prompt_token_ids", + empty_is_missing=not empty_prompt_is_exact, + ) token_ids = _exact_token_ids( choice_data.get("token_ids"), field="Completions token_ids" ) @@ -235,29 +558,237 @@ def _completion_tokens( if not complete_pairs: pair_ids = [] pair_logprobs = [ - float(value) if isinstance(value, (int, float)) else math.nan + float(value) + if isinstance(value, (int, float)) and not isinstance(value, bool) + else math.nan for value in logprobs.get("token_logprobs") or [] ] + pair_includes_prompt = ( + echo + and prompt_ids is not None + and token_ids is not None + and pair_ids == [*prompt_ids, *token_ids] + ) + token_ids_include_prompt = ( + echo + and prompt_ids is not None + and token_ids is not None + and bool(pair_ids) + and token_ids == [*prompt_ids, *pair_ids] + ) if token_ids is not None and pair_ids and token_ids != pair_ids: - raise ValueError("Response token IDs disagree with completion logprobs") + if not pair_includes_prompt and not token_ids_include_prompt: + raise ValueError("Response token IDs disagree with completion logprobs") selected = token_ids if token_ids is not None else pair_ids or None - if selected is not None and len(pair_logprobs) != len(selected): - pair_logprobs = [math.nan] * len(selected) - return prompt_ids, selected, pair_logprobs + prompt_logprobs: list[float] = [] + completion_logprobs = pair_logprobs + if echo and prompt_ids is None and token_ids is None: + # A combined logprobs.tokens carrier does not reveal where an echoed + # prompt ends. Let the string history tokenize prompt and completion + # independently rather than mislabeling the prompt as sampled output. + selected = None + completion_logprobs = [] + elif echo and prompt_ids is not None and selected is not None: + if pair_includes_prompt: + prompt_logprobs = pair_logprobs[: len(prompt_ids)] + completion_logprobs = pair_logprobs[len(prompt_ids) :] + elif len(pair_logprobs) == len(prompt_ids) + len(selected): + prompt_logprobs = pair_logprobs[: len(prompt_ids)] + completion_logprobs = pair_logprobs[len(prompt_ids) :] + elif token_ids_include_prompt: + selected = selected[len(prompt_ids) :] + completion_logprobs = pair_logprobs + elif selected[: len(prompt_ids)] == prompt_ids and ( + len(tokens) == len(selected) + and all(isinstance(value, str) for value in tokens) + and "".join(cast(str, value) for value in tokens) == choice.text + ): + if pair_logprobs and len(pair_logprobs) != len(selected): + raise ValueError("Completions token IDs and logprobs differ in length") + prompt_logprobs = pair_logprobs[: len(prompt_ids)] + selected = selected[len(prompt_ids) :] + completion_logprobs = pair_logprobs[len(prompt_ids) :] + elif selected[: len(prompt_ids)] == prompt_ids: + selected = None + completion_logprobs = [] + elif pair_logprobs and len(pair_logprobs) != len(selected): + raise ValueError("Completions token IDs and logprobs differ in length") + elif selected is not None and pair_logprobs and len(pair_logprobs) != len(selected): + raise ValueError("Completions token IDs and logprobs differ in length") + if selected is not None and not completion_logprobs: + completion_logprobs = [math.nan] * len(selected) + return prompt_ids, selected, prompt_logprobs, completion_logprobs + + +def _completion_tokens( + response: Completion, + *, + echo: bool = False, + empty_prompt_is_exact: bool = False, +) -> tuple[list[int] | None, list[int] | None, list[float]]: + prompt, completion, _, completion_logprobs = _completion_evidence( + response, + echo=echo, + empty_prompt_is_exact=empty_prompt_is_exact, + ) + return prompt, completion, completion_logprobs + + +@dataclass(frozen=True) +class _ResponseGeneration: + prompt_token_ids: list[int] | None + output_token_ids: list[int] | None + output_logprobs: list[float] + output_indices: list[int] + output_text: str | None + + +def _responses_output_is_sampled(item: object) -> bool: + return not str(_field(item, "type", "")).endswith("_output") + + +def _response_generations(response: Response) -> list[_ResponseGeneration]: + data = _dump(response) + raw_generations = data.get("token_generations") + if raw_generations is None: + return [] + if not isinstance(raw_generations, list): + raise ValueError("Responses token_generations exact metadata must be a list") + if not raw_generations: + raise ValueError( + "Responses token_generations must be omitted when exact evidence " + "is unavailable" + ) + generations: list[_ResponseGeneration] = [] + assigned_outputs: set[int] = set() + last_output_index = -1 + for index, raw_generation in enumerate(raw_generations): + generation = _string_dict(raw_generation) + if generation is None: + raise ValueError( + f"Responses token_generations[{index}] must be a JSON object" + ) + prompt = _exact_token_ids( + generation.get("prompt_token_ids"), + field=f"Responses token_generations[{index}].prompt_token_ids", + empty_is_missing=True, + ) + if prompt is None: + raise ValueError( + f"Responses token_generations[{index}].prompt_token_ids must " + "contain the full non-empty generation prompt" + ) + raw_indices = generation.get("output_indices") + if not isinstance(raw_indices, list) or any( + not isinstance(value, int) or isinstance(value, bool) or value < 0 + for value in raw_indices + ): + raise ValueError( + f"Responses token_generations[{index}].output_indices must be " + "a list of non-negative integers" + ) + output_indices = [ + value + for value in raw_indices + if isinstance(value, int) and not isinstance(value, bool) + ] + if len(set(output_indices)) != len(output_indices): + raise ValueError( + f"Responses token_generations[{index}].output_indices contains duplicates" + ) + if output_indices != sorted(output_indices): + raise ValueError( + f"Responses token_generations[{index}].output_indices must be ordered" + ) + if output_indices and output_indices[0] <= last_output_index: + raise ValueError( + "Responses token_generations output_indices must be ordered " + "and nonoverlapping" + ) + if output_indices: + last_output_index = output_indices[-1] + if any(value >= len(response.output) for value in output_indices): + raise ValueError( + f"Responses token_generations[{index}].output_indices is out of bounds" + ) + if assigned_outputs.intersection(output_indices): + raise ValueError( + "Responses token_generations output_indices overlap between generations" + ) + assigned_outputs.update(output_indices) + output = generation.get("output_tokens") + output_ids: list[int] | None + output_logprobs: list[float] + if output is None: + output_ids, output_logprobs = None, [] + output_text = None + else: + output_ids, output_logprobs = _pairs( + output, + require_token_ids=True, + field=f"Responses token_generations[{index}].output_tokens", + ) + if not output_ids: + output_ids = None + output_logprobs = [] + output_texts = [ + item.get("text") + for item in (_string_dict(value) for value in output) + if item is not None + ] + output_text = ( + "".join(cast(str, text) for text in output_texts) + if len(output_texts) == len(output) + and all(isinstance(text, str) for text in output_texts) + else None + ) + if output_ids is None and any( + _responses_output_is_sampled(response.output[value]) + for value in output_indices + ): + raise ValueError( + f"Responses token_generations[{index}].output_tokens must contain " + "exact tokens for its sampled output items" + ) + generations.append( + _ResponseGeneration( + prompt_token_ids=prompt, + output_token_ids=output_ids, + output_logprobs=output_logprobs, + output_indices=output_indices, + output_text=output_text, + ) + ) + required_outputs = { + index + for index, item in enumerate(response.output) + if _responses_output_is_sampled(item) + } + if not required_outputs.issubset(assigned_outputs): + raise ValueError( + "Responses token_generations output_indices must cover every sampled " + "output item" + ) + return generations def _responses_tokens( response: Response, -) -> tuple[None, list[int] | None, list[float]]: +) -> tuple[list[int] | None, list[int] | None, list[float]]: data = _dump(response) - if "raw_output_tokens" in data: - token_ids, logprobs = _pairs( - data["raw_output_tokens"], - require_token_ids=True, - field="Responses raw_output_tokens", - ) - if token_ids or not data.get("output"): - return None, token_ids, logprobs + generations = _response_generations(response) + if generations: + if len(generations) != 1: + raise ValueError( + "A multi-generation Responses exchange must be tokenized through " + "its protocol history" + ) + generation = generations[0] + return ( + generation.prompt_token_ids, + generation.output_token_ids, + generation.output_logprobs, + ) token_ids: list[int] = [] logprobs: list[float] = [] saw_rendered_output = False @@ -288,18 +819,27 @@ def _responses_tokens( def _messages_tokens( response: Message, -) -> tuple[None, list[int] | None, list[float]]: +) -> tuple[list[int] | None, list[int] | None, list[float]]: data = _dump(response) + prompt_ids = _exact_token_ids( + data.get("prompt_token_ids"), + field="Messages prompt_token_ids", + empty_is_missing=True, + ) token_ids = _exact_token_ids(data.get("token_ids"), field="Messages token_ids") if token_ids == [] and data.get("content"): token_ids = None logprobs = [ - float(value) if isinstance(value, (int, float)) else math.nan + float(value) + if isinstance(value, (int, float)) and not isinstance(value, bool) + else math.nan for value in data.get("logprobs") or [] ] - if token_ids is None or len(logprobs) != len(token_ids): + if token_ids is not None and logprobs and len(logprobs) != len(token_ids): + raise ValueError("Messages token IDs and logprobs differ in length") + if token_ids is None or not logprobs: logprobs = [math.nan] * len(token_ids or []) - return None, token_ids, logprobs + return prompt_ids, token_ids, logprobs def _exchange_list(trajectory: Trajectory, model: str | None) -> list[Exchange]: @@ -310,7 +850,16 @@ def _exchange_list(trajectory: Trajectory, model: str | None) -> list[Exchange]: *trajectory.exchanges.messages, ] if model is not None: - exchanges = [exchange for exchange in exchanges if exchange.model == model] + has_exact_match = any(exchange.model == model for exchange in exchanges) + exchanges = [ + exchange + for exchange in exchanges + if ( + exchange.model == model + if has_exact_match + else _model_matches(exchange.model, model) + ) + ] if not exchanges: raise ValueError(f"Trajectory contains no exchanges for model {model!r}") models = {exchange.model for exchange in exchanges} @@ -325,22 +874,59 @@ def _exchange_list(trajectory: Trajectory, model: str | None) -> list[Exchange]: ) -def _artifact_config(model: str) -> _TokenizerConfig: +def _artifact_name(model: str) -> str: + return model.removeprefix("wandb-artifact:///") + + +def _artifact_identity(model: str) -> str: + """Return an alias/version-independent checkpoint identity for base-model cache.""" + + path = _artifact_name(model) + name = path.rsplit("/", 1)[-1] + return path[: -len(name)] + name.split(":", 1)[0] + + +_ARTIFACT_BASE_MODELS: dict[str, str] = {} + + +def _artifact_base_model(identity: str) -> str: + if cached := _ARTIFACT_BASE_MODELS.get(identity): + return cached from wandb.apis.public import Api - artifact_path = model.removeprefix("wandb-artifact:///") + artifact_path = identity if ":" not in artifact_path.rsplit("/", 1)[-1]: artifact_path = f"{artifact_path}:latest" artifact = Api().artifact(artifact_path) metadata = artifact.metadata base_model = metadata.get("base_model") or metadata.get("wandb.base_model") if not isinstance(base_model, str): - raise ValueError(f"Checkpoint {model!r} does not identify its base model") + raise ValueError(f"Checkpoint {identity!r} does not identify its base model") + if len(_ARTIFACT_BASE_MODELS) >= 1024: + _ARTIFACT_BASE_MODELS.pop(next(iter(_ARTIFACT_BASE_MODELS))) + _ARTIFACT_BASE_MODELS[identity] = base_model + return base_model + + +def _artifact_config(model: str) -> _TokenizerConfig: + from wandb.apis.public import Api + + artifact_path = _artifact_name(model) + if ":" not in artifact_path.rsplit("/", 1)[-1]: + artifact_path = f"{artifact_path}:latest" + artifact = Api().artifact(artifact_path) + metadata = artifact.metadata + base_model = metadata.get("base_model") or metadata.get("wandb.base_model") + identity = _artifact_identity(model) + if isinstance(base_model, str): + if len(_ARTIFACT_BASE_MODELS) >= 1024 and identity not in _ARTIFACT_BASE_MODELS: + _ARTIFACT_BASE_MODELS.pop(next(iter(_ARTIFACT_BASE_MODELS))) + _ARTIFACT_BASE_MODELS[identity] = base_model renderer = metadata.get("renderer") renderer = renderer if isinstance(renderer, dict) else {} kwargs = renderer.get("chat_template_kwargs") return _TokenizerConfig( - base_model=base_model, + base_model=_artifact_base_model(identity), revision=( renderer.get("tokenizer_revision") if isinstance(renderer.get("tokenizer_revision"), str) @@ -368,7 +954,8 @@ def _tokenizer_config(model: str, base_model: str | None) -> _TokenizerConfig: return _TokenizerConfig(model) -def _load_tokenizer(config: _TokenizerConfig) -> _Tokenizer: +@lru_cache(maxsize=8) +def _cached_tokenizer(base_model: str, revision: str | None) -> Tokenizer: try: from transformers import AutoTokenizer except ImportError as exc: @@ -376,18 +963,25 @@ def _load_tokenizer(config: _TokenizerConfig) -> _Tokenizer: "Tokenizer fallback requires ART's backend or tinker dependencies" ) from exc try: - return _as_tokenizer( - AutoTokenizer.from_pretrained( - config.base_model, - revision=config.revision, - ) + tokenizer = AutoTokenizer.from_pretrained( + base_model, + revision=revision, ) + if base_model.startswith("deepseek-ai/DeepSeek-V4-"): + from ..megatron.dsv4.tokenizer import get_dsv4_tokenizer + + tokenizer = get_dsv4_tokenizer(cast("PreTrainedTokenizerBase", tokenizer)) + return _as_tokenizer(tokenizer) except Exception as exc: raise ValueError( - f"Could not load tokenizer for {config.base_model!r}; pass base_model explicitly" + f"Could not load tokenizer for {base_model!r}; pass base_model explicitly" ) from exc +def _load_tokenizer(config: _TokenizerConfig) -> Tokenizer: + return _cached_tokenizer(config.base_model, config.revision) + + def _ids(value: object) -> list[int]: if (input_ids := getattr(value, "input_ids", None)) is not None: value = input_ids @@ -745,7 +1339,7 @@ def _response_message( def _template_ids( - tokenizer: _Tokenizer, + tokenizer: Tokenizer, exchange: Exchange, *, completed: bool, @@ -803,7 +1397,11 @@ def _exchange_tokens( if isinstance(exchange, ChatCompletionsExchange): return _chat_tokens(exchange.response) if isinstance(exchange, CompletionsExchange): - return _completion_tokens(exchange.response) + return _completion_tokens( + exchange.response, + echo=exchange.request.get("echo") is True, + empty_prompt_is_exact=exchange.request.get("prompt") in ("", []), + ) if isinstance(exchange, ResponsesExchange): return _responses_tokens(exchange.response) if isinstance(exchange, MessagesExchange): @@ -811,20 +1409,32 @@ def _exchange_tokens( raise TypeError(f"Unknown exchange type: {type(exchange)!r}") -def _visible_logprobs(exchange: Exchange) -> list[tuple[str, float]]: +def _visible_logprobs( + exchange: Exchange, *, source: object | None = None +) -> list[tuple[str, float]]: values: list[tuple[str, float]] = [] if isinstance(exchange, ChatCompletionsExchange): - logprobs = exchange.response.choices[0].logprobs - entries = (logprobs.content or logprobs.refusal or []) if logprobs else [] - for entry in entries: + choice = ( + _chat_choice(source) if source is not None else exchange.response.choices[0] + ) + entries = _chat_logprob_entries(choice) + decoder = codecs.getincrementaldecoder("utf-8")() + for index, entry in enumerate(entries): data = _dump(entry) raw_bytes = data.get("bytes") if isinstance(raw_bytes, list): try: - text = bytes(raw_bytes).decode("utf-8") + next_data = ( + _dump(entries[index + 1]) if index + 1 < len(entries) else {} + ) + text = decoder.decode( + bytes(raw_bytes), + final=not isinstance(next_data.get("bytes"), list), + ) except (TypeError, ValueError, UnicodeDecodeError): return [] else: + decoder = codecs.getincrementaldecoder("utf-8")() text = data.get("token") logprob = data.get("logprob") if isinstance(text, str) and isinstance(logprob, (int, float)): @@ -838,7 +1448,16 @@ def _visible_logprobs(exchange: Exchange) -> list[tuple[str, float]]: if logprob is not None: values.append((text, float(logprob))) elif isinstance(exchange, ResponsesExchange): - for output in _dump(exchange.response).get("output") or []: + outputs = exchange.response.output + if source is not None: + selected = _responses_source_outputs(source) + if selected is None: + return [] + selected_exchange, output_indices = selected + if selected_exchange is not exchange: + raise ValueError("Responses source belongs to a different exchange") + outputs = [exchange.response.output[index] for index in output_indices] + for output in outputs: for content in _dump(output).get("content") or []: for entry in _dump(content).get("logprobs") or []: data = _dump(entry) @@ -849,20 +1468,184 @@ def _visible_logprobs(exchange: Exchange) -> list[tuple[str, float]]: return values -def _align_visible_logprobs( - tokenizer: _Tokenizer | None, completion: list[int], exchange: Exchange -) -> list[float] | None: - values = _visible_logprobs(exchange) +def _sampled_text(exchange: Exchange, *, source: object | None = None) -> str | None: + visible = "".join(text for text, _ in _visible_logprobs(exchange, source=source)) + if visible: + return visible + if isinstance(exchange, ChatCompletionsExchange): + choice = ( + _chat_choice(source) if source is not None else exchange.response.choices[0] + ) + content = choice.message.content + return content if isinstance(content, str) else None + if isinstance(exchange, CompletionsExchange): + return exchange.response.choices[0].text + if isinstance(exchange, MessagesExchange): + parts = [ + block.text + for block in exchange.response.content + if isinstance(block, TextBlock) + ] + return "".join(parts) if parts else None + if isinstance(exchange, ResponsesExchange): + outputs = exchange.response.output + if source is not None: + selected = _responses_source_outputs(source) + if selected is None: + return None + selected_exchange, output_indices = selected + if selected_exchange is not exchange: + raise ValueError("Responses source belongs to a different exchange") + outputs = [exchange.response.output[index] for index in output_indices] + parts: list[str] = [] + for output in outputs: + data = _dump(output) + if data.get("type") == "message": + parts.append(_responses_output_text(data.get("content"))) + return "".join(parts) if parts else None + return None + + +def _preserve_sampled_prefix( + prompt: list[int], + canonical_prefix: list[int], + sampled_outputs: list[_SampledOutput], + tokenizer: Tokenizer, +) -> list[int] | None: + repaired = prompt + for sampled in sampled_outputs: + if sampled.text is None: + return None + rendered_ids = _ids(tokenizer(sampled.text, add_special_tokens=False)) + if rendered_ids == sampled.token_ids: + continue + start = sampled.start + if ( + repaired[:start] != canonical_prefix[:start] + or repaired[start : start + len(rendered_ids)] != rendered_ids + ): + return None + repaired = [ + *repaired[:start], + *sampled.token_ids, + *repaired[start + len(rendered_ids) :], + ] + return repaired if repaired[: len(canonical_prefix)] == canonical_prefix else None + + +def _warn_prefix_retokenization() -> None: + global _WARNED_PREFIX_RETOKENIZATION + if _WARNED_PREFIX_RETOKENIZATION: + return + _WARNED_PREFIX_RETOKENIZATION = True + warnings.warn( + "Inference prompt token IDs retokenized an earlier sampled response; ART " + "preserved the original sampled token IDs and logprobs. Prefer a service " + "with prefix token-ID preservation, such as Caladan.", + stacklevel=3, + ) + + +def _retained_output_suffix( + *, + prompt: Sequence[int], + output: Sequence[int], + logprobs: Sequence[float], + later_prompt: Sequence[int], +) -> tuple[list[int], list[float]] | None: + if list(later_prompt[: len(prompt)]) != list(prompt): + return None + continuation = later_prompt[len(prompt) :] + for start in range(len(output)): + suffix = list(output[start:]) + if list(continuation[: len(suffix)]) == suffix: + return ( + suffix, + list(logprobs[start:]) + if len(logprobs) == len(output) + else [math.nan] * len(suffix), + ) + return None + + +def _visible_token_evidence( + tokenizer: Tokenizer | None, + exchange: Exchange, + *, + source: object | None = None, + sampled_text: str | None = None, +) -> tuple[list[int], list[float]] | None: + values = _visible_logprobs(exchange, source=source) if not values or tokenizer is None: return None + if sampled_text is not None: + if "".join(text for text, _ in values) != sampled_text: + matches: list[list[tuple[str, float]]] = [] + for start in range(len(values)): + combined = "" + for end in range(start, len(values)): + combined += values[end][0] + if combined == sampled_text: + matches.append(values[start : end + 1]) + break + if len(combined) >= len(sampled_text): + break + if len(matches) != 1: + return None + values = matches[0] + logprobs = [logprob for _, logprob in values] + text = "".join(text for text, _ in values) + if any(not value for value, _ in values): + token_ids = _ids(tokenizer(text, add_special_tokens=False)) + return (token_ids, logprobs) if len(token_ids) == len(values) else None + try: + contextual = cast(_OffsetTokenizer, tokenizer)( + text, + add_special_tokens=False, + return_offsets_mapping=True, + ) + except (TypeError, ValueError, NotImplementedError): + contextual = None + contextual_data = _string_dict(contextual) + offsets = ( + contextual_data.get("offset_mapping") if contextual_data is not None else None + ) + if isinstance(offsets, list) and len(offsets) == len(values): + boundaries: list[tuple[int, int]] = [] + cursor = 0 + for value, _ in values: + boundaries.append((cursor, cursor + len(value))) + cursor += len(value) + if offsets == boundaries: + token_ids = _ids(contextual) + if len(token_ids) == len(values): + return token_ids, logprobs token_ids: list[int] = [] - logprobs: list[float] = [] - for text, logprob in values: + for text, _ in values: encoded = _ids(tokenizer(text, add_special_tokens=False)) if len(encoded) != 1: return None token_ids.append(encoded[0]) - logprobs.append(logprob) + return token_ids, logprobs + + +def _align_visible_logprobs( + tokenizer: Tokenizer | None, + completion: list[int], + exchange: Exchange, + *, + source: object | None = None, + sampled_text: str | None = None, +) -> list[float] | None: + evidence = _visible_token_evidence( + tokenizer, + exchange, + source=source, + sampled_text=sampled_text, + ) + if evidence is None: + return None + token_ids, logprobs = evidence left: list[int] = [] cursor = 0 @@ -895,18 +1678,14 @@ def _align_visible_logprobs( def _legacy_tokenize( - trajectory: Trajectory, - base_model: str | None, + history: LegacyHistory, *, - chat_template: str | None, - chat_template_kwargs: Mapping[str, object] | None, -) -> TokenizedTrajectory: - if trajectory.additional_histories: - raise ValueError("Tokenization requires one history") + model: str, +) -> TokenizedHistory: token_ids: list[int] = [] logprobs: list[float] = [] flags: list[TokenFlag] = [] - for item in trajectory.messages_and_choices: + for item in history.messages_and_choices: if not isinstance(item, Choice): continue prompt, completion, completion_logprobs = _chat_choice_tokens(item, {}) @@ -932,23 +1711,24 @@ def _legacy_tokenize( flags.extend([TokenFlag.EXACT | TokenFlag.SAMPLED] * len(completion)) if not token_ids: raise ValueError("Trajectory contains no trainable choices") - return TokenizedTrajectory( + return TokenizedHistory( + model=model, token_ids=token_ids, logprobs=logprobs, flags=flags, - underlying=trajectory, ) -def tokenize_one( +def _tokenize_exchange_trajectory( trajectory: Trajectory, base_model: str | None, *, model: str | None, chat_template: str | None, chat_template_kwargs: Mapping[str, object] | None, - tokenizer_instance: _Tokenizer | None = None, -) -> TokenizedTrajectory: + tokenizer_instance: Tokenizer | None = None, + _trace: _TraceBuilder | None = None, +) -> TokenizedHistory: if trajectory.exchanges and ( trajectory.messages_and_choices or trajectory.tools is not None @@ -958,11 +1738,16 @@ def tokenize_one( "A trajectory cannot contain both exchanges and legacy histories" ) if not trajectory.exchanges: + if trajectory.additional_histories: + raise ValueError("Tokenization requires one history") + if model is None: + raise ValueError("Legacy trajectory tokenization requires model=") return _legacy_tokenize( - trajectory, - base_model, - chat_template=chat_template, - chat_template_kwargs=chat_template_kwargs, + LegacyHistory( + messages_and_choices=trajectory.messages_and_choices, + tools=trajectory.tools, + ), + model=model, ) exchanges = _exchange_list(trajectory, model) selected_model = exchanges[0].model @@ -983,9 +1768,13 @@ def tokenize_one( token_ids: list[int] = [] logprobs: list[float] = [] flags: list[TokenFlag] = [] + source_keys: list[_SampledSourceKey | None] = [] + sources: dict[_SampledSourceKey, object] = {} response_histories: dict[ str, tuple[list[dict[str, Any]] | None, ResponsesExchange] ] = {} + sampled_outputs: list[_SampledOutput] = [] + previous_render_state: tuple[Exchange, list[dict[str, Any]] | None] | None = None def fallback_config() -> _TokenizerConfig: nonlocal config @@ -1016,6 +1805,10 @@ def fallback_config() -> _TokenizerConfig: messages_override: list[dict[str, Any]] | None = None if isinstance(exchange, ResponsesExchange): request = exchange.request + if request.get("conversation") is not None and prompt is None: + raise ValueError( + "Responses conversation history requires exact prompt tokens" + ) try: messages_override = _responses_messages(request) except ValueError: @@ -1023,21 +1816,25 @@ def fallback_config() -> _TokenizerConfig: raise previous = request.get("previous_response_id") if previous is not None: - if not isinstance(previous, str) or previous not in response_histories: + if not isinstance(previous, str): + raise ValueError("Responses previous_response_id must be text") + if previous not in response_histories and prompt is None: raise ValueError( - "Responses exchange refers to a previous response outside this trajectory" + "Responses exchange refers to a previous response outside this " + "trajectory without exact prompt tokens" ) - previous_messages, previous_exchange = response_histories[previous] - if prompt is None: - if previous_messages is None or messages_override is None: - raise ValueError( - "Responses history cannot be rendered without exact prompt tokens" - ) - messages_override = [ - *previous_messages, - _response_message(previous_exchange), - *messages_override, - ] + if previous in response_histories: + previous_messages, previous_exchange = response_histories[previous] + if prompt is None: + if previous_messages is None or messages_override is None: + raise ValueError( + "Responses history cannot be rendered without exact prompt tokens" + ) + messages_override = [ + *previous_messages, + _response_message(previous_exchange), + *messages_override, + ] response_histories[exchange.response.id] = (messages_override, exchange) if prompt is None: resolved_config = fallback_config() @@ -1056,6 +1853,19 @@ def fallback_config() -> _TokenizerConfig: resolved_config = fallback_config() if tokenizer is None: tokenizer = _load_tokenizer(resolved_config) + rendered_prompt = ( + _template_ids( + tokenizer, + exchange, + completed=False, + config=resolved_config, + chat_template=chat_template, + chat_template_kwargs=chat_template_kwargs, + messages_override=messages_override, + ) + if prompt_is_exact + else prompt + ) completed = _template_ids( tokenizer, exchange, @@ -1065,11 +1875,11 @@ def fallback_config() -> _TokenizerConfig: chat_template_kwargs=chat_template_kwargs, messages_override=messages_override, ) - if completed[: len(prompt)] != prompt: + if completed[: len(rendered_prompt)] != rendered_prompt: raise ValueError( "Completed response does not extend its generation prompt" ) - completion = completed[len(prompt) :] + completion = completed[len(rendered_prompt) :] completion_logprobs = _align_visible_logprobs( tokenizer, completion, exchange ) or [math.nan] * len(completion) @@ -1079,19 +1889,76 @@ def fallback_config() -> _TokenizerConfig: flags.extend( [TokenFlag.EXACT if prompt_is_exact else TokenFlag(0)] * len(prompt) ) + source_keys.extend([None] * len(prompt)) elif len(prompt) < len(token_ids) or prompt[: len(token_ids)] != token_ids: - raise ValueError( - "Exchanges do not resolve to one append-only token history" + resolved_config = fallback_config() + if tokenizer is None: + tokenizer = _load_tokenizer(resolved_config) + repaired = _preserve_sampled_prefix( + prompt, + token_ids, + sampled_outputs, + tokenizer, ) - else: + if repaired is None: + if previous_render_state is None: + raise ValueError( + "Inference prompts do not form one append-only history" + ) + current_render = _template_ids( + tokenizer, + exchange, + completed=False, + config=resolved_config, + chat_template=chat_template, + chat_template_kwargs=chat_template_kwargs, + messages_override=messages_override, + ) + previous_exchange, previous_messages = previous_render_state + previous_render = _template_ids( + tokenizer, + previous_exchange, + completed=True, + config=resolved_config, + chat_template=chat_template, + chat_template_kwargs=chat_template_kwargs, + messages_override=previous_messages, + ) + previous_canonical = _preserve_sampled_prefix( + previous_render, + token_ids, + sampled_outputs, + tokenizer, + ) + if ( + previous_canonical is None + or current_render[: len(previous_render)] != previous_render + ): + raise ValueError( + "Rendered inference prompts do not form one append-only history" + ) + repaired = [ + *previous_canonical, + *current_render[len(previous_render) :], + ] + prompt = repaired + prompt_is_exact = False + _warn_prefix_retokenization() suffix = prompt[len(token_ids) :] token_ids.extend(suffix) logprobs.extend([math.nan] * len(suffix)) - if prompt_is_exact: + flags.extend([TokenFlag(0)] * len(suffix)) + source_keys.extend([None] * len(suffix)) + else: + suffix = prompt[len(token_ids) :] + token_ids.extend(suffix) + logprobs.extend([math.nan] * len(suffix)) + if prompt_is_exact: flags = [flag | TokenFlag.EXACT for flag in flags] flags.extend( [TokenFlag.EXACT if prompt_is_exact else TokenFlag(0)] * len(suffix) ) + source_keys.extend([None] * len(suffix)) if len(completion_logprobs) != len(completion): completion_logprobs = _align_visible_logprobs( tokenizer, completion, exchange @@ -1102,10 +1969,2927 @@ def fallback_config() -> _TokenizerConfig: if completion_is_exact: completion_flag |= TokenFlag.EXACT flags.extend([completion_flag] * len(completion)) + source_key = _exchange_sampled_source_key(exchange) + source_keys.extend([source_key] * len(completion)) + sources[source_key] = exchange + if completion_is_exact: + sampled_outputs.append( + _SampledOutput( + text=_sampled_text(exchange), + token_ids=list(completion), + start=len(token_ids) - len(completion), + ) + ) + previous_render_state = (exchange, messages_override) - return TokenizedTrajectory( + tokenized = TokenizedHistory( + model=selected_model, + token_ids=token_ids, + logprobs=logprobs, + flags=flags, + ) + if _trace is not None: + _trace.set(tokenized, source_keys, sources) + return tokenized + + +def _unique_exchanges(history: History) -> list[Exchange]: + sources: Sequence[object] + if isinstance(history, ChatCompletionsHistory): + if len(history.messages) != len(history.message_sources): + raise ValueError("messages and message_sources differ in length") + sources = history.message_sources + elif isinstance(history, AnthropicMessagesHistory): + if len(history.messages) != len(history.message_sources): + raise ValueError("messages and message_sources differ in length") + sources = history.message_sources + elif isinstance(history, ResponsesHistory): + if len(history.input) != len(history.input_sources): + raise ValueError("input and input_sources differ in length") + sources = history.input_sources + else: + raise TypeError(f"Unsupported history type: {type(history).__name__}") + + exchanges: list[Exchange] = [] + seen: set[int] = set() + for source in sources: + exchange = getattr(source, "exchange", None) + if not isinstance( + exchange, + ( + ChatCompletionsExchange, + CompletionsExchange, + ResponsesExchange, + MessagesExchange, + ), + ): + continue + identity = id(exchange) + if identity in seen: + continue + seen.add(identity) + if isinstance(exchange, ChatCompletionsExchange): + choice_indexes = { + getattr(item, "choice_index", None) + for item in sources + if getattr(item, "exchange", None) is exchange + and getattr(item, "choice_index", None) is not None + } + if choice_indexes: + exchange = exchange.model_copy( + update={ + "response": exchange.response.model_copy( + update={ + "choices": [ + choice + for choice in exchange.response.choices + if choice.index in choice_indexes + ] + } + ) + } + ) + exchanges.append(exchange) + return sorted(exchanges, key=lambda item: (item.start_time, item.end_time)) + + +def _without_reasoning(message: Mapping[str, object]) -> dict[str, object]: + visible = dict(message) + visible.pop("reasoning", None) + visible.pop("reasoning_content", None) + return visible + + +def _chat_choice_message(source: object) -> dict[str, Any] | None: + exchange = getattr(source, "exchange", None) + choice_index = getattr(source, "choice_index", None) + if not isinstance(exchange, ChatCompletionsExchange) or not isinstance( + choice_index, int + ): + return None + choice = next( + (item for item in exchange.response.choices if item.index == choice_index), + None, + ) + return ( + choice.message.model_dump(mode="python", exclude_none=True) + if choice is not None + else None + ) + + +def _source_index(value: object, *, length: int, field: str) -> int: + if not isinstance(value, int) or isinstance(value, bool) or not 0 <= value < length: + raise ValueError(f"{field} is out of bounds") + return value + + +def _chat_output_indices(source: object) -> tuple[int, ...] | None: + value = getattr(source, "output_indices", None) + if value is None: + return None + if not isinstance(value, tuple) or any( + not isinstance(index, int) or isinstance(index, bool) for index in value + ): + raise ValueError("Chat output source indices are invalid") + if tuple(sorted(set(value))) != value: + raise ValueError("Chat output source indices are not strictly ordered") + return value + + +def _chat_choice(source: object) -> Choice: + exchange = getattr(source, "exchange", None) + choice_index = getattr(source, "choice_index", None) + if not isinstance(exchange, ChatCompletionsExchange): + raise ValueError("Chat choice source has the wrong exchange type") + if ( + not isinstance(choice_index, int) + or isinstance(choice_index, bool) + or choice_index < 0 + ): + raise ValueError("Chat choice source index is invalid") + choice = next( + (item for item in exchange.response.choices if item.index == choice_index), + None, + ) + if choice is None: + raise ValueError("Chat choice source index is out of bounds") + return choice + + +def _validate_history_sources(history: History) -> None: + if isinstance(history, ChatCompletionsHistory): + from ._history import normalize_chat_message + + if len(history.messages) != len(history.message_sources): + raise ValueError("messages and message_sources differ in length") + for message, source in zip( + history.messages, history.message_sources, strict=True + ): + if source is None: + continue + exchange = source.exchange + if exchange.model != history.model: + raise ValueError( + "Chat Completions history model no longer matches its source " + "exchange" + ) + output_indices = _chat_output_indices(source) + expected: list[dict[str, Any]] = [] + if isinstance(exchange, ChatCompletionsExchange): + if source.choice_index is not None: + if ( + source.request_index is not None + or output_indices is not None + or source.generation_index is not None + ): + raise ValueError("Chat source has conflicting indices") + choice_message = _chat_choice(source).message.model_dump( + mode="python", exclude_none=True + ) + expected.append(choice_message) + visible = _without_reasoning(choice_message) + if visible != choice_message: + expected.append(visible) + elif source.request_index is not None: + if ( + output_indices is not None + or source.generation_index is not None + ): + raise ValueError("Chat source has conflicting indices") + request_messages = exchange.request.get("messages", []) + request_index = _source_index( + source.request_index, + length=len(request_messages), + field="Chat request source index", + ) + expected.append(dict(request_messages[request_index])) + else: + raise ValueError("Chat source has no request or choice index") + elif isinstance(exchange, MessagesExchange): + if ( + source.choice_index is not None + or source.generation_index is not None + ): + raise ValueError("Anthropic-to-Chat source has invalid indices") + if source.request_index is not None: + if output_indices is not None: + raise ValueError("Anthropic-to-Chat source has invalid indices") + request_messages = exchange.request.get("messages", []) + request_index = _source_index( + source.request_index, + length=len(request_messages), + field="Anthropic request source index", + ) + expected.extend( + _anthropic_messages( + {"messages": [request_messages[request_index]]} + ) + ) + elif output_indices is None: + expected.extend( + _anthropic_messages( + { + "system": exchange.request.get("system"), + "messages": [], + } + ) + ) + elif output_indices == (0,): + expected.append(_response_message(exchange)) + visible_blocks = [ + block.model_dump(mode="python", exclude_none=True) + for block in exchange.response.content + if getattr(block, "type", None) + not in {"thinking", "redacted_thinking"} + ] + expected.extend( + _anthropic_messages( + { + "messages": [ + {"role": "assistant", "content": visible_blocks} + ] + } + ) + ) + else: + raise ValueError("Anthropic response source index is out of bounds") + elif isinstance(exchange, ResponsesExchange): + if source.request_index is not None: + if ( + source.choice_index is not None + or output_indices is not None + or source.generation_index is not None + ): + raise ValueError("Responses-to-Chat source has invalid indices") + request_input = exchange.request.get("input") + if isinstance(request_input, list): + request_index = _source_index( + source.request_index, + length=len(request_input), + field="Responses request source index", + ) + for end in range(request_index + 1, len(request_input) + 1): + projected = _responses_messages( + {"input": request_input[request_index:end]} + ) + if len(projected) == 1: + expected.extend(projected) + elif isinstance(request_input, str): + _source_index( + source.request_index, + length=1, + field="Responses request source index", + ) + expected.extend(_responses_messages({"input": request_input})) + else: + raise ValueError( + "Responses request source index is out of bounds" + ) + else: + if source.choice_index is not None: + raise ValueError("Responses-to-Chat source has invalid indices") + if output_indices is not None: + resolved_indices = tuple( + _source_index( + output_index, + length=len(exchange.response.output), + field="Responses output source index", + ) + for output_index in output_indices + ) + if source.generation_index is not None: + generations = _response_generations(exchange.response) + generation_index = _source_index( + source.generation_index, + length=len(generations), + field="Responses generation source index", + ) + generation_outputs = generations[ + generation_index + ].output_indices + if not resolved_indices and generation_outputs: + raise ValueError( + "Responses empty output source references a " + "generation with output items" + ) + if any( + output_index not in generation_outputs + for output_index in resolved_indices + ): + raise ValueError( + "Responses output source does not belong to " + "its generation" + ) + if not resolved_indices and source.generation_index is not None: + expected.append({"role": "assistant", "content": ""}) + expected.extend( + _responses_messages( + { + "input": [ + exchange.response.output[ + output_index + ].model_dump(mode="python", exclude_none=True) + for output_index in resolved_indices + ] + } + ) + ) + elif message.get("role") == "system": + expected.extend( + _responses_messages( + {"instructions": exchange.request.get("instructions")} + ) + ) + elif source.generation_index is not None: + expected.append({"role": "assistant", "content": ""}) + actual = normalize_chat_message(message) + if not any( + actual == normalize_chat_message(candidate) for candidate in expected + ): + raise ValueError( + "Chat Completions history no longer matches its source exchange" + ) + return + if isinstance(history, AnthropicMessagesHistory): + from ._history import _anthropic_message_key + + if len(history.messages) != len(history.message_sources): + raise ValueError("messages and message_sources differ in length") + for message, source in zip( + history.messages, history.message_sources, strict=True + ): + if source is None: + continue + if source.exchange.model != history.model: + raise ValueError( + "Anthropic Messages history model no longer matches its source " + "exchange" + ) + if source.request_index is None: + expected = { + "role": "assistant", + "content": [ + block.model_dump(mode="json", exclude_none=True) + for block in source.exchange.response.content + ], + } + else: + request_messages = source.exchange.request.get("messages", []) + request_index = _source_index( + source.request_index, + length=len(request_messages), + field="Anthropic request source index", + ) + expected = request_messages[request_index] + matches = expected is not None and message == expected + if not matches and source.request_index is None and expected is not None: + matches = _anthropic_message_key(message) == _anthropic_message_key( + cast(MessageParam, expected), visible_only=True + ) + if not matches: + raise ValueError( + "Anthropic Messages history no longer matches its source exchange" + ) + return + if isinstance(history, ResponsesHistory): + from ._history import _responses_input + + if len(history.input) != len(history.input_sources): + raise ValueError("input and input_sources differ in length") + + for item, source in zip(history.input, history.input_sources, strict=True): + if source is None: + continue + if source.exchange.model != history.model: + raise ValueError( + "Responses history model no longer matches its source exchange" + ) + if source.request_index is not None and source.output_index is not None: + raise ValueError("Responses source has conflicting indices") + if source.output_index is not None: + output = source.exchange.response.output + output_index = _source_index( + source.output_index, + length=len(output), + field="Responses output source index", + ) + expected = output[output_index].model_dump( + mode="json", exclude_none=True + ) + if source.generation_index is not None: + generations = _response_generations(source.exchange.response) + generation_index = _source_index( + source.generation_index, + length=len(generations), + field="Responses generation source index", + ) + if output_index not in generations[generation_index].output_indices: + raise ValueError( + "Responses output source does not belong to its generation" + ) + elif source.request_index is not None: + request_input = _responses_input(source.exchange.request.get("input")) + request_index = _source_index( + source.request_index, + length=len(request_input), + field="Responses request source index", + ) + if source.generation_index is not None: + raise ValueError("Responses request source has a generation index") + expected = request_input[request_index] + else: + generations = _response_generations(source.exchange.response) + if source.generation_index is None: + raise ValueError( + "Responses source has no request, output, or generation index" + ) + generation_index = _source_index( + source.generation_index, + length=len(generations), + field="Responses generation source index", + ) + if ( + generation_index != len(generations) - 1 + or generations[generation_index].output_indices + ): + raise ValueError( + "Responses generation-only source must refer to a terminal " + "generation without native output items" + ) + expected = {"role": "assistant", "content": ""} + if expected is None or item != expected: + raise ValueError( + "Responses history no longer matches its source exchange" + ) + + +def _trajectory_from_history(history: History) -> Trajectory: + _validate_history_sources(history) + exchanges = _unique_exchanges(history) + if not exchanges: + raise ValueError( + "History has no source exchanges; local history-only rendering is not yet possible" + ) + from . import TrajectoryExchanges + + return Trajectory( + exchanges=TrajectoryExchanges( + chat_completions=[ + item for item in exchanges if isinstance(item, ChatCompletionsExchange) + ], + completions=[ + item for item in exchanges if isinstance(item, CompletionsExchange) + ], + responses=[ + item for item in exchanges if isinstance(item, ResponsesExchange) + ], + messages=[item for item in exchanges if isinstance(item, MessagesExchange)], + ) + ) + + +def _last_source_exchange(sources: Sequence[object]) -> Exchange | None: + for source in reversed(sources): + exchange = getattr(source, "exchange", None) + if isinstance( + exchange, + ( + ChatCompletionsExchange, + CompletionsExchange, + ResponsesExchange, + MessagesExchange, + ), + ): + return exchange + return None + + +@dataclass(frozen=True) +class _HistoryRenderState: + needs_render: bool + context_changed: bool = False + projection_matches: bool | None = None + + +def _matches_final_chat_exchange(history: ChatCompletionsHistory) -> bool | None: + from ._history import normalize_chat_message + + for source in reversed(history.message_sources): + if ( + source is None + or not isinstance(source.exchange, ChatCompletionsExchange) + or source.choice_index is None + ): + continue + choice = _chat_choice(source) + expected = [ + *source.exchange.request.get("messages", []), + choice.message.model_dump(mode="python", exclude_none=True), + ] + return [normalize_chat_message(message) for message in history.messages] == [ + normalize_chat_message(message) for message in expected + ] + return None + + +def _history_render_state(history: History) -> _HistoryRenderState: + if isinstance(history, ChatCompletionsHistory): + if any(source is None for source in history.message_sources): + return _HistoryRenderState(needs_render=True, projection_matches=False) + for message, source in zip( + history.messages, history.message_sources, strict=True + ): + if source is None or (original := _chat_choice_message(source)) is None: + continue + if dict(message) != original and dict(message) == _without_reasoning( + original + ): + return _HistoryRenderState(needs_render=True) + exchange = _last_source_exchange(history.message_sources) + if not isinstance(exchange, ChatCompletionsExchange): + return _HistoryRenderState(needs_render=True) + context_changed = ( + history.tools != exchange.request.get("tools") + or history.chat_template != exchange.request.get("chat_template") + or history.chat_template_kwargs + != exchange.request.get("chat_template_kwargs") + ) + if context_changed: + return _HistoryRenderState(needs_render=True, context_changed=True) + projection_matches = _matches_final_chat_exchange(history) + if projection_matches is None: + projection_matches = _history_matches_projection(history) + return _HistoryRenderState( + needs_render=not projection_matches, + projection_matches=projection_matches, + ) + if isinstance(history, AnthropicMessagesHistory): + if any(source is None for source in history.message_sources): + return _HistoryRenderState(needs_render=True, projection_matches=False) + for message, source in zip( + history.messages, history.message_sources, strict=True + ): + if source is None or source.request_index is not None: + continue + expected = { + "role": "assistant", + "content": [ + block.model_dump(mode="json", exclude_none=True) + for block in source.exchange.response.content + ], + } + if message != expected: + return _HistoryRenderState(needs_render=True) + exchange = _last_source_exchange(history.message_sources) + if not isinstance(exchange, MessagesExchange): + return _HistoryRenderState(needs_render=False) + context_changed = ( + history.system != exchange.request.get("system") + or history.tools != exchange.request.get("tools") + or history.chat_template != exchange.request.get("chat_template") + or history.chat_template_kwargs + != exchange.request.get("chat_template_kwargs") + ) + if context_changed: + return _HistoryRenderState(needs_render=True, context_changed=True) + projection_matches = _history_matches_projection(history) + return _HistoryRenderState( + needs_render=not projection_matches, + projection_matches=projection_matches, + ) + if isinstance(history, ResponsesHistory): + if any(source is None for source in history.input_sources): + return _HistoryRenderState(needs_render=True, projection_matches=False) + exchange = _last_source_exchange(history.input_sources) + if not isinstance(exchange, ResponsesExchange): + return _HistoryRenderState(needs_render=False) + context_changed = ( + history.instructions != exchange.request.get("instructions") + or history.tools != exchange.request.get("tools") + or history.chat_template != exchange.request.get("chat_template") + or history.chat_template_kwargs + != exchange.request.get("chat_template_kwargs") + ) + if context_changed: + return _HistoryRenderState(needs_render=True, context_changed=True) + projection_matches = _history_matches_projection(history) + return _HistoryRenderState( + needs_render=not projection_matches, + projection_matches=projection_matches, + ) + return _HistoryRenderState(needs_render=False) + + +def _source_signature(source: object) -> tuple[object, ...] | None: + if source is None: + return None + exchange = getattr(source, "exchange", None) + if not isinstance( + exchange, + ( + ChatCompletionsExchange, + CompletionsExchange, + ResponsesExchange, + MessagesExchange, + ), + ): + return None + response_id = getattr(exchange.response, "id", None) + choice_index = getattr(source, "choice_index", None) + generation_index = getattr(source, "generation_index", None) + evidence_fingerprint: str | None = None + if isinstance(exchange, ChatCompletionsExchange) and isinstance(choice_index, int): + evidence_fingerprint = _sampled_evidence_fingerprint( + exchange, protocol="chat_completions", index=choice_index + ) + elif isinstance(exchange, ResponsesExchange) and isinstance(generation_index, int): + evidence_fingerprint = _sampled_evidence_fingerprint( + exchange, protocol="responses", index=generation_index + ) + elif isinstance(exchange, MessagesExchange) and ( + getattr(source, "output_index", None) == 0 + or _chat_output_indices(source) == (0,) + ): + evidence_fingerprint = _sampled_evidence_fingerprint( + exchange, protocol="messages", index=0 + ) + return ( + type(source), + type(exchange), + exchange.start_time, + exchange.end_time, + response_id, + evidence_fingerprint, + getattr(source, "request_index", None), + choice_index, + getattr(source, "output_index", None), + _chat_output_indices(source), + generation_index, + getattr(source, "prompt_index", None), + ) + + +def _sources_match(left: Sequence[object], right: Sequence[object]) -> bool: + return [_source_signature(item) for item in left] == [ + _source_signature(item) for item in right + ] + + +def _history_matches_projection(history: History) -> bool: + exchanges = _unique_exchanges(history) + if not exchanges: + return False + from . import TrajectoryExchanges + from ._history import ( + anthropic_messages_histories, + chat_completions_histories, + responses_histories, + ) + + trajectory = Trajectory( + exchanges=TrajectoryExchanges( + chat_completions=[ + item for item in exchanges if isinstance(item, ChatCompletionsExchange) + ], + completions=[ + item for item in exchanges if isinstance(item, CompletionsExchange) + ], + responses=[ + item for item in exchanges if isinstance(item, ResponsesExchange) + ], + messages=[item for item in exchanges if isinstance(item, MessagesExchange)], + ) + ) + if isinstance(history, ChatCompletionsHistory): + try: + candidates: Sequence[History] = chat_completions_histories( + trajectory, model=history.model + ) + except ValueError as error: + if "no Chat Completions exchanges" not in str(error): + raise + if trajectory.exchanges.messages and not trajectory.exchanges.responses: + candidates = [ + candidate.as_chat_completions_history() + for candidate in anthropic_messages_histories( + trajectory, model=history.model + ) + ] + elif trajectory.exchanges.responses and not trajectory.exchanges.messages: + candidates = [ + candidate.as_chat_completions_history() + for candidate in responses_histories( + trajectory, model=history.model + ) + ] + else: + return False + return any( + isinstance(candidate, ChatCompletionsHistory) + and candidate.messages == history.messages + and candidate.tools == history.tools + and candidate.chat_template == history.chat_template + and candidate.chat_template_kwargs == history.chat_template_kwargs + and _sources_match(candidate.message_sources, history.message_sources) + for candidate in candidates + ) + if isinstance(history, AnthropicMessagesHistory): + candidates = anthropic_messages_histories(trajectory, model=history.model) + return any( + isinstance(candidate, AnthropicMessagesHistory) + and candidate.messages == history.messages + and candidate.system == history.system + and candidate.tools == history.tools + and candidate.chat_template == history.chat_template + and candidate.chat_template_kwargs == history.chat_template_kwargs + and _sources_match(candidate.message_sources, history.message_sources) + for candidate in candidates + ) + if isinstance(history, ResponsesHistory): + candidates = responses_histories(trajectory, model=history.model) + return any( + isinstance(candidate, ResponsesHistory) + and candidate.input == history.input + and candidate.instructions == history.instructions + and candidate.tools == history.tools + and candidate.conversation == history.conversation + and candidate.previous_response_id == history.previous_response_id + and candidate.chat_template == history.chat_template + and candidate.chat_template_kwargs == history.chat_template_kwargs + and _sources_match(candidate.input_sources, history.input_sources) + for candidate in candidates + ) + return False + + +def _response_generation_text( + response: Response, generation: _ResponseGeneration +) -> str | None: + parts: list[str] = [] + for output_index in generation.output_indices: + item = _dump(response.output[output_index]) + kind = item.get("type") + if kind == "message": + parts.append(_responses_output_text(item.get("content"))) + elif kind == "reasoning": + parts.append(_responses_reasoning_text(item)) + else: + return None + return "".join(parts) or None + + +def _tokenize_exact_responses_history( + history: ResponsesHistory, + *, + base_model: str | None, + tokenizer: Tokenizer | None, + _trace: _TraceBuilder | None = None, +) -> TokenizedHistory | None: + generation_keys: list[tuple[ResponsesExchange, int]] = [] + retained_output_indices: dict[tuple[int, int], set[int]] = {} + seen: set[tuple[int, int]] = set() + for source in history.input_sources: + if source is None or source.generation_index is None: + continue + key = (id(source.exchange), source.generation_index) + if source.output_index is not None: + retained_output_indices.setdefault(key, set()).add(source.output_index) + if key not in seen: + seen.add(key) + generation_keys.append((source.exchange, source.generation_index)) + if not generation_keys: + return None + + token_ids: list[int] = [] + logprobs: list[float] = [] + flags: list[TokenFlag] = [] + source_keys: list[_SampledSourceKey | None] = [] + sources: dict[_SampledSourceKey, object] = {} + sampled_outputs: list[_SampledOutput] = [] + for position, (exchange, generation_index) in enumerate(generation_keys): + generations = _response_generations(exchange.response) + if not 0 <= generation_index < len(generations): + raise ValueError("Responses source generation index is out of bounds") + generation = generations[generation_index] + prompt = generation.prompt_token_ids + output = generation.output_token_ids + if prompt is None or output is None: + return None + retained = retained_output_indices.get((id(exchange), generation_index), set()) + if retained != set(generation.output_indices): + if position + 1 >= len(generation_keys): + return None + next_exchange, next_generation_index = generation_keys[position + 1] + next_generations = _response_generations(next_exchange.response) + if not 0 <= next_generation_index < len(next_generations): + raise ValueError("Responses source generation index is out of bounds") + next_prompt = next_generations[next_generation_index].prompt_token_ids + if next_prompt is None: + return None + retained_suffix = _retained_output_suffix( + prompt=prompt, + output=output, + logprobs=generation.output_logprobs, + later_prompt=next_prompt, + ) + if retained_suffix is None: + return None + output, output_logprobs = retained_suffix + output_text = None + else: + output_logprobs = generation.output_logprobs + output_text = generation.output_text or _response_generation_text( + exchange.response, generation + ) + if not token_ids: + token_ids.extend(prompt) + logprobs.extend([math.nan] * len(prompt)) + flags.extend([TokenFlag.EXACT] * len(prompt)) + source_keys.extend([None] * len(prompt)) + elif prompt[: len(token_ids)] == token_ids: + suffix = prompt[len(token_ids) :] + token_ids.extend(suffix) + logprobs.extend([math.nan] * len(suffix)) + flags = [flag | TokenFlag.EXACT for flag in flags] + flags.extend([TokenFlag.EXACT] * len(suffix)) + source_keys.extend([None] * len(suffix)) + else: + if tokenizer is None and sampled_outputs: + tokenizer = _load_tokenizer( + _tokenizer_config(history.model, base_model) + ) + repaired = ( + _preserve_sampled_prefix(prompt, token_ids, sampled_outputs, tokenizer) + if tokenizer is not None + else None + ) + if repaired is None: + raise ValueError( + "Responses token generations do not form one append-only history" + ) + _warn_prefix_retokenization() + suffix = repaired[len(token_ids) :] + token_ids.extend(suffix) + logprobs.extend([math.nan] * len(suffix)) + flags.extend([TokenFlag.EXACT] * len(suffix)) + source_keys.extend([None] * len(suffix)) + token_ids.extend(output) + logprobs.extend(output_logprobs) + flags.extend([TokenFlag.EXACT | TokenFlag.SAMPLED] * len(output)) + source = next( + ( + item + for item in history.input_sources + if item is not None + and item.exchange is exchange + and item.generation_index == generation_index + ), + None, + ) + if source is None: + raise AssertionError("Responses generation has no history source") + source_key = _sampled_source_key(source) + source_keys.extend([source_key] * len(output)) + sources[source_key] = source + sampled_outputs.append( + _SampledOutput( + text=output_text, + token_ids=list(output), + start=len(token_ids) - len(output), + ) + ) + tokenized = TokenizedHistory( + model=history.model, token_ids=token_ids, logprobs=logprobs, flags=flags, - underlying=trajectory, + ) + if _trace is not None: + _trace.set(tokenized, source_keys, sources) + return tokenized + + +def _responses_source_outputs( + source: object, +) -> tuple[ResponsesExchange, tuple[int, ...]] | None: + exchange = getattr(source, "exchange", None) + if not isinstance(exchange, ResponsesExchange): + return None + output_indices = _chat_output_indices(source) + if output_indices is None: + return None + return exchange, tuple( + _source_index( + index, + length=len(exchange.response.output), + field="Responses source output index", + ) + for index in output_indices + ) + + +def _responses_source_generation( + source: object, +) -> tuple[ResponsesExchange, _ResponseGeneration, tuple[int, ...]] | None: + selected = _responses_source_outputs(source) + generation_index = getattr(source, "generation_index", None) + if selected is None or generation_index is None: + return None + if not isinstance(generation_index, int) or isinstance(generation_index, bool): + raise ValueError("Responses source generation index is invalid") + exchange, output_indices = selected + generations = _response_generations(exchange.response) + if not 0 <= generation_index < len(generations): + raise ValueError("Responses source generation index is out of bounds") + generation = generations[generation_index] + if bool(output_indices) != bool(generation.output_indices): + raise ValueError("Responses empty output source does not match its generation") + if any(index not in generation.output_indices for index in output_indices): + raise ValueError( + "Responses source output index does not belong to its generation" + ) + return exchange, generation, output_indices + + +def _responses_generation_messages(source: object) -> list[dict[str, Any]] | None: + selected = _responses_source_generation(source) + if selected is None: + return None + exchange, _, output_indices = selected + return _responses_messages( + { + "input": [ + exchange.response.output[index].model_dump( + mode="python", exclude_none=True + ) + for index in output_indices + ] + } + ) + + +def _responses_generation_full_tokens( + source: object, +) -> tuple[list[int] | None, list[float]]: + selected = _responses_source_generation(source) + if selected is None: + return None, [] + _, generation, output_indices = selected + if output_indices != tuple(generation.output_indices): + return None, [] + messages = _responses_generation_messages(source) + if messages is None or (messages and len(messages) != 1): + return None, [] + return generation.output_token_ids, generation.output_logprobs + + +def _chat_source_full_tokens( + source: object, +) -> tuple[list[int] | None, list[float]]: + exchange = getattr(source, "exchange", None) + if isinstance(exchange, ChatCompletionsExchange): + choice_index = getattr(source, "choice_index", None) + if choice_index is None: + return None, [] + choice = next( + item for item in exchange.response.choices if item.index == choice_index + ) + return _chat_choice_output_tokens(choice) + if isinstance(exchange, ResponsesExchange): + return _responses_generation_full_tokens(source) + if isinstance(exchange, MessagesExchange): + if _chat_output_indices(source) != (0,): + return None, [] + _, tokens, logprobs = _exchange_tokens(exchange) + return tokens, logprobs + return None, [] + + +def _chat_source_prompt_tokens(source: object) -> list[int] | None: + exchange = getattr(source, "exchange", None) + if isinstance(exchange, ChatCompletionsExchange): + choice_index = getattr(source, "choice_index", None) + if choice_index is None: + return None + choice = next( + item for item in exchange.response.choices if item.index == choice_index + ) + prompt, _, _ = _chat_choice_tokens(choice, _dump(exchange.response)) + return prompt + if isinstance(exchange, ResponsesExchange): + generation_index = getattr(source, "generation_index", None) + generations = _response_generations(exchange.response) + if isinstance(generation_index, int): + if not 0 <= generation_index < len(generations): + raise ValueError("Responses source generation index is out of bounds") + return generations[generation_index].prompt_token_ids + return _responses_tokens(exchange.response)[0] + if isinstance(exchange, MessagesExchange): + if _chat_output_indices(source) != (0,): + return None + return _messages_tokens(exchange.response)[0] + return None + + +def _source_is_sampled(source: object) -> bool: + exchange = getattr(source, "exchange", None) + if isinstance(exchange, ChatCompletionsExchange): + return getattr(source, "choice_index", None) is not None + if isinstance(exchange, MessagesExchange): + return _chat_output_indices(source) == (0,) + if isinstance(exchange, ResponsesExchange): + if getattr(source, "generation_index", None) is not None: + return True + selected = _responses_source_outputs(source) + return selected is not None and any( + _responses_output_is_sampled(selected[0].response.output[index]) + for index in selected[1] + ) + return False + + +def _tokenize_exact_projected_chat_history( + history: ChatCompletionsHistory, + *, + projection_validated: bool = False, + _trace: _TraceBuilder | None = None, +) -> TokenizedHistory | None: + if not projection_validated and not _history_matches_projection(history): + return None + sampled_sources: list[object] = [] + seen: set[tuple[object, ...]] = set() + for message, source in zip(history.messages, history.message_sources, strict=True): + signature = _source_signature(source) + if ( + message.get("role") == "assistant" + and source is not None + and signature is not None + and signature not in seen + and _source_is_sampled(source) + ): + seen.add(signature) + sampled_sources.append(source) + if not sampled_sources: + return None + + final_source = sampled_sources[-1] + final_prompt = _chat_source_prompt_tokens(final_source) + final_output, final_logprobs = _chat_source_full_tokens(final_source) + if final_prompt is None or final_output is None: + return None + + token_ids = [*final_prompt, *final_output] + logprobs = [ + *([math.nan] * len(final_prompt)), + *( + final_logprobs + if len(final_logprobs) == len(final_output) + else [math.nan] * len(final_output) + ), + ] + flags = [ + *([TokenFlag.EXACT] * len(final_prompt)), + *([TokenFlag.EXACT | TokenFlag.SAMPLED] * len(final_output)), + ] + final_key = _sampled_source_key(final_source) + source_keys: list[_SampledSourceKey | None] = [ + *([None] * len(final_prompt)), + *([final_key] * len(final_output)), + ] + sources: dict[_SampledSourceKey, object] = {final_key: final_source} + for index, source in enumerate(sampled_sources[:-1]): + prompt = _chat_source_prompt_tokens(source) + output, output_logprobs = _chat_source_full_tokens(source) + if ( + prompt is None + or output is None + or list(final_prompt[: len(prompt)]) != prompt + ): + return None + retained = next( + ( + evidence + for later_source in sampled_sources[index + 1 :] + if (later_prompt := _chat_source_prompt_tokens(later_source)) + is not None + and ( + evidence := _retained_output_suffix( + prompt=prompt, + output=output, + logprobs=output_logprobs, + later_prompt=later_prompt, + ) + ) + is not None + ), + None, + ) + if retained is None: + return None + retained_ids, retained_logprobs = retained + start = len(prompt) + end = start + len(retained_ids) + if final_prompt[start:end] != retained_ids: + return None + flags[start:end] = [TokenFlag.EXACT | TokenFlag.SAMPLED] * len(retained_ids) + logprobs[start:end] = retained_logprobs + source_key = _sampled_source_key(source) + source_keys[start:end] = [source_key] * len(retained_ids) + sources[source_key] = source + if history.model is None: + raise ValueError("History tokenization requires a model") + tokenized = TokenizedHistory( + model=history.model, + token_ids=token_ids, + logprobs=logprobs, + flags=flags, + ) + if _trace is not None: + _trace.set(tokenized, source_keys, sources) + return tokenized + + +def _chat_message_parts(message: Mapping[str, object]) -> list[tuple[str, str]]: + parts: list[tuple[str, str]] = [] + reasoning = message.get("reasoning") or message.get("reasoning_content") + if isinstance(reasoning, str) and reasoning: + parts.append(("reasoning", reasoning)) + if content := _content_text(message.get("content")): + parts.append(("content", content)) + refusal = message.get("refusal") + if isinstance(refusal, str) and refusal: + parts.append(("content", refusal)) + tool_calls = message.get("tool_calls") + if isinstance(tool_calls, list): + for tool_call in tool_calls: + call = _string_dict(tool_call) + function = _string_dict(call.get("function")) if call is not None else None + if function is None: + continue + for value in (function.get("name"), function.get("arguments")): + if isinstance(value, str) and value: + parts.append(("tool_call", value)) + return parts + + +def _chat_message_text_slot_groups( + message: dict[str, Any], +) -> list[list[tuple[dict[str, Any], str]]]: + groups: list[list[tuple[dict[str, Any], str]]] = [] + reasoning_key = ( + "reasoning" + if isinstance(message.get("reasoning"), str) and message["reasoning"] + else "reasoning_content" + ) + if isinstance(message.get(reasoning_key), str) and message[reasoning_key]: + groups.append([(message, reasoning_key)]) + content = message.get("content") + if isinstance(content, str) and content: + groups.append([(message, "content")]) + elif isinstance(content, list): + content_slots = [ + (block, "text") + for block in content + if ( + isinstance(block, dict) + and block.get("type") in {"input_text", "output_text", "text"} + and isinstance(block.get("text"), str) + and block["text"] + ) + ] + if content_slots: + groups.append(content_slots) + if isinstance(message.get("refusal"), str) and message["refusal"]: + groups.append([(message, "refusal")]) + tool_calls = message.get("tool_calls") + if isinstance(tool_calls, list): + for call in tool_calls: + if not isinstance(call, dict): + continue + function = call.get("function") + if not isinstance(function, dict): + continue + for key in ("name", "arguments"): + if isinstance(function.get(key), str) and function[key]: + groups.append([(function, key)]) + return groups + + +def _chat_source_tokens( + source: object, + text: str, + *, + part: str, + full_tokens: tuple[list[int] | None, list[float]] | None = None, +) -> tuple[list[int] | None, list[float]]: + exchange = getattr(source, "exchange", None) + if isinstance(exchange, ChatCompletionsExchange): + tokens, logprobs = ( + full_tokens if full_tokens is not None else _chat_source_full_tokens(source) + ) + if tokens is None: + return None, [] + sampled_text = _sampled_text(exchange, source=source) + message = _dump( + next( + item + for item in exchange.response.choices + if item.index == getattr(source, "choice_index", None) + ).message + ) + if ( + part == "content" + and sampled_text == text + and not any(message.get(key) for key in ("reasoning", "tool_calls")) + ): + return tokens, logprobs + return None, [] + if isinstance(exchange, MessagesExchange) and _chat_output_indices(source) == (0,): + block_type = "thinking" if part == "reasoning" else "text" + blocks = [ + block + for block in exchange.response.content + if getattr(block, "type", None) == block_type + ] + block_text = "".join( + str(getattr(block, "thinking" if part == "reasoning" else "text", "")) + for block in blocks + ) + if block_text == text: + token_ids: list[int] = [] + logprobs: list[float] = [] + for block in blocks: + extra = block.model_extra or {} + block_ids = _exact_token_ids( + extra.get("token_ids"), field="Messages content token_ids" + ) + block_logprobs = extra.get("logprobs") + if block_ids is None or not isinstance(block_logprobs, list): + break + if len(block_ids) != len(block_logprobs): + raise ValueError( + "Messages content token IDs and logprobs differ in length" + ) + token_ids.extend(block_ids) + logprobs.extend(float(value) for value in block_logprobs) + else: + return token_ids, logprobs + if part != "content" or any( + getattr(block, "type", None) in {"thinking", "redacted_thinking"} + for block in exchange.response.content + ): + return None, [] + _, tokens, logprobs = _messages_tokens(exchange.response) + return tokens, logprobs + if isinstance(exchange, ResponsesExchange) and _source_is_sampled(source): + selected = _responses_source_outputs(source) + output_indices = selected[1] if selected is not None else () + if len(output_indices) == 1: + output_index = output_indices[0] + item = _dump(exchange.response.output[output_index]) + if item.get("type") == "message" and part == "content": + item_text = _responses_output_text(item.get("content")) + if item_text == text: + token_ids: list[int] = [] + logprobs: list[float] = [] + for content in item.get("content") or []: + pairs, pair_logprobs = _pairs( + _dump(content).get("logprobs"), + field="Responses content logprobs", + ) + if not pairs: + break + token_ids.extend(pairs) + logprobs.extend(pair_logprobs) + else: + return token_ids, logprobs + # Aggregate generation evidence cannot be partitioned safely across + # multiple projected messages. Item-local pairs above remain usable; + # otherwise render this message without claiming exact token identity. + return None, [] + return None, [] + + +def _source_covers_complete_sampled_message( + message: Mapping[str, object], source: object +) -> bool: + from ._history import normalize_chat_message + + exchange = getattr(source, "exchange", None) + if isinstance(exchange, ChatCompletionsExchange): + expected = _chat_choice_message(source) + return expected is not None and normalize_chat_message( + message + ) == normalize_chat_message(expected) + if isinstance(exchange, MessagesExchange): + if _chat_output_indices(source) != (0,): + return False + return normalize_chat_message(message) == normalize_chat_message( + _response_message(exchange) + ) + if not isinstance(exchange, ResponsesExchange): + return False + generation_index = getattr(source, "generation_index", None) + if not isinstance(generation_index, int): + return False + generations = _response_generations(exchange.response) + if not 0 <= generation_index < len(generations): + raise ValueError("Responses source generation index is out of bounds") + output_indices = _chat_output_indices(source) + if output_indices is None: + return False + if not output_indices: + return message.get("role") == "assistant" and not _chat_message_parts(message) + generation = generations[generation_index] + if any(index not in generation.output_indices for index in output_indices): + raise ValueError( + "Responses source output index does not belong to its generation" + ) + projected = _responses_generation_messages(source) + if projected is None: + return False + return len(projected) == 1 and normalize_chat_message( + message + ) == normalize_chat_message(projected[0]) + + +def _tokenize_chat_view( + history: ChatCompletionsHistory, + *, + base_model: str | None, + tokenizer: Tokenizer | None, + chat_template: str | None, + chat_template_kwargs: Mapping[str, object] | None, + _projection_matches: bool | None = None, + _trace: _TraceBuilder | None = None, +) -> TokenizedHistory: + _validate_history_sources(history) + config = ( + _TokenizerConfig(base_model or history.model or "") + if tokenizer is not None + or (base_model is not None and chat_template is not None) + else _tokenizer_config(history.model or "", base_model) + ) + if tokenizer is None: + if not history.model and base_model is None: + raise ValueError("History tokenization requires a model or base_model") + tokenizer = _load_tokenizer(config) + assert tokenizer is not None + resolved_tokenizer = tokenizer + messages = [dict(message) for message in history.messages] + kwargs = { + **(config.chat_template_kwargs or {}), + **(history.chat_template_kwargs or {}), + **(chat_template_kwargs or {}), + } + last_exchange = _last_source_exchange(history.message_sources) + if isinstance(last_exchange, MessagesExchange) and isinstance( + thinking := last_exchange.request.get("thinking"), dict + ): + kwargs.setdefault("enable_thinking", thinking.get("type") == "enabled") + if budget := thinking.get("budget_tokens"): + kwargs.setdefault("thinking_budget", budget) + template = chat_template or history.chat_template or config.chat_template + ends_with_assistant = bool(messages) and messages[-1].get("role") == "assistant" + + def render( + selected_messages: list[dict[str, Any]], *, add_generation_prompt: bool + ) -> list[int]: + return _ids( + resolved_tokenizer.apply_chat_template( + selected_messages, + tools=history.tools, + tokenize=True, + add_generation_prompt=add_generation_prompt, + **({"chat_template": template} if template is not None else {}), + **kwargs, + ) + ) + + rendered = render(messages, add_generation_prompt=not ends_with_assistant) + if any( + message.get("role") == "assistant" + and isinstance(message.get("reasoning"), str) + and message["reasoning"] + and not message.get("reasoning_content") + for message in messages + ): + without_reasoning = deepcopy(messages) + for message in without_reasoning: + message.pop("reasoning", None) + if ( + render( + without_reasoning, + add_generation_prompt=not ends_with_assistant, + ) + == rendered + ): + messages = deepcopy(messages) + for message in messages: + reasoning = message.pop("reasoning", None) + if isinstance(reasoning, str) and reasoning: + message.setdefault("reasoning_content", reasoning) + rendered = render( + messages, + add_generation_prompt=not ends_with_assistant, + ) + if any( + message.get("role") == "assistant" + and isinstance(message.get("refusal"), str) + and message["refusal"] + for message in messages + ): + without_refusals = deepcopy(messages) + for message in without_refusals: + message.pop("refusal", None) + if ( + render( + without_refusals, + add_generation_prompt=not ends_with_assistant, + ) + == rendered + ): + messages = deepcopy(messages) + for message in messages: + refusal = message.pop("refusal", None) + if not isinstance(refusal, str) or not refusal: + continue + content = message.get("content") + if isinstance(content, str): + message["content"] = content + refusal + elif isinstance(content, list): + message["content"] = [ + *content, + {"type": "text", "text": refusal}, + ] + elif content is None: + message["content"] = refusal + else: + raise ValueError( + "Cannot render an assistant refusal with this content shape" + ) + rendered = render( + messages, + add_generation_prompt=not ends_with_assistant, + ) + + prompt_cache: dict[int, list[int] | None] = {} + output_cache: dict[int, tuple[list[int] | None, list[float]]] = {} + + def source_prompt_tokens(source: object) -> list[int] | None: + key = id(source) + if key not in prompt_cache: + prompt_cache[key] = _chat_source_prompt_tokens(source) + return prompt_cache[key] + + def source_output_tokens( + source: object, + ) -> tuple[list[int] | None, list[float]]: + key = id(source) + if key not in output_cache: + output_cache[key] = _chat_source_full_tokens(source) + return output_cache[key] + + def source_matches_context(source: object) -> bool: + exchange = getattr(source, "exchange", None) + return ( + isinstance(exchange, ChatCompletionsExchange) + and history.tools == exchange.request.get("tools") + and history.chat_template == exchange.request.get("chat_template") + and history.chat_template_kwargs + == exchange.request.get("chat_template_kwargs") + ) + + exact_prefix_length = 0 + if ( + chat_template is None + and chat_template_kwargs is None + and _projection_matches is True + ): + for message_index, (message, source) in enumerate( + zip(history.messages, history.message_sources, strict=True) + ): + if message.get("role") != "assistant" or source is None: + continue + source_prompt = source_prompt_tokens(source) + if source_prompt and source_matches_context(source): + rendered_prompt = render( + messages[:message_index], add_generation_prompt=True + ) + if rendered[: len(rendered_prompt)] == rendered_prompt: + rendered = [ + *source_prompt, + *rendered[len(rendered_prompt) :], + ] + exact_prefix_length = len(source_prompt) + break + + positions_by_first_token: dict[int, list[int]] = {} + for index, token_id in enumerate(rendered): + positions_by_first_token.setdefault(token_id, []).append(index) + locations_by_needle: dict[tuple[int, ...], list[tuple[int, int]]] = {} + + def locations(needle: Sequence[int], start: int) -> list[tuple[int, int]]: + if not needle: + return [] + key = tuple(needle) + if key not in locations_by_needle: + locations_by_needle[key] = [ + (index, index + len(key)) + for index in positions_by_first_token.get(key[0], []) + if rendered[index : index + len(key)] == list(key) + ] + spans = locations_by_needle[key] + return spans[bisect_left(spans, (start, -1)) :] + + replacements: list[ + tuple[ + int, + int, + list[int], + list[float], + bool, + _SampledSourceKey, + object, + ] + ] = [] + search_cursor = 0 + sampled_texts = { + text + for message, source in zip(messages, history.message_sources, strict=True) + if message.get("role") == "assistant" + and source is not None + and _source_is_sampled(source) + for _, text in _chat_message_parts(message) + } + part_ids_cache: dict[str, list[int]] = {} + + def part_ids(text: str) -> list[int]: + if text not in part_ids_cache: + part_ids_cache[text] = _ids( + resolved_tokenizer(text, add_special_tokens=False) + ) + return part_ids_cache[text] + + direct_render: list[int] = [] + direct_bounds: list[tuple[int, int]] = [] + for message in messages: + start = len(direct_render) + for _, text in _chat_message_parts(message): + direct_render.extend(part_ids(text)) + direct_bounds.append((start, len(direct_render))) + if direct_render != rendered and not ( + exact_prefix_length + and len(direct_render) == len(rendered) + and direct_render[exact_prefix_length:] == rendered[exact_prefix_length:] + ): + direct_bounds = [] + + def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: + prefix = 0 + while ( + prefix < len(rendered) + and prefix < len(probe) + and rendered[prefix] == probe[prefix] + ): + prefix += 1 + suffix = 0 + while ( + suffix < len(rendered) - prefix + and suffix < len(probe) - prefix + and rendered[-suffix - 1] == probe[-suffix - 1] + ): + suffix += 1 + end = len(rendered) - suffix + return (prefix, end) if prefix < end else None + + marked_bounds: dict[int, tuple[int, int]] = {} + marked_part_bounds: dict[int, list[tuple[int, int]]] = {} + if not direct_bounds: + marked_messages = deepcopy(messages) + marker_prefix = f"ART_TRAJECTORY_{id(marked_messages):x}_" + markers: dict[str, tuple[int, int, Literal["start", "end"]]] = {} + part_counts: dict[int, int] = {} + part_whitespace: dict[tuple[int, int], tuple[str, str]] = {} + for message_index, (message, source) in enumerate( + zip(marked_messages, history.message_sources, strict=True) + ): + if ( + message.get("role") != "assistant" + or source is None + or not _source_is_sampled(source) + ): + continue + slot_groups = _chat_message_text_slot_groups(message) + if not slot_groups: + continue + part_counts[message_index] = len(slot_groups) + for part_index, slots in enumerate(slot_groups): + start = f"{marker_prefix}{message_index}_{part_index}_START" + end = f"{marker_prefix}{message_index}_{part_index}_END" + first, first_key = slots[0] + last, last_key = slots[-1] + first_text = str(first[first_key]) + last_text = str(last[last_key]) + leading = first_text[: len(first_text) - len(first_text.lstrip())] + trailing = last_text[len(last_text.rstrip()) :] + whitespace_only = ( + first is last and first_key == last_key and not first_text.strip() + ) + if whitespace_only: + trailing = "" + if first is last and first_key == last_key: + core = first_text[ + len(leading) : len(first_text) - len(trailing) + if trailing + else len(first_text) + ] + first[first_key] = leading + start + core + end + trailing + else: + first[first_key] = leading + start + first_text[len(leading) :] + last[last_key] = ( + last_text[: len(last_text) - len(trailing)] + end + trailing + ) + part_whitespace[(message_index, part_index)] = ( + ("", "") if whitespace_only else (leading, trailing) + ) + markers[start] = (message_index, part_index, "start") + markers[end] = (message_index, part_index, "end") + if markers: + try: + marked_text = resolved_tokenizer.apply_chat_template( + marked_messages, + tools=history.tools, + tokenize=False, + add_generation_prompt=not ends_with_assistant, + **({"chat_template": template} if template is not None else {}), + **kwargs, + ) + except (TypeError, NotImplementedError): + marked_text = None + if isinstance(marked_text, str): + marker_pattern = re.compile( + rf"{re.escape(marker_prefix)}\d+_\d+_(?:START|END)" + ) + matches = list(marker_pattern.finditer(marked_text)) + found_markers = [match.group(0) for match in matches] + else: + matches = [] + found_markers = [] + if ( + isinstance(marked_text, str) + and len(found_markers) == len(markers) + and set(found_markers) == set(markers) + ): + unmarked_parts: list[str] = [] + char_bounds: dict[tuple[int, int], list[int]] = {} + source_cursor = 0 + target_cursor = 0 + for match in matches: + position = match.start() + marker = match.group(0) + message_index, part_index, boundary = markers[marker] + chunk = marked_text[source_cursor:position] + unmarked_parts.append(chunk) + target_cursor += len(chunk) + char_bounds.setdefault((message_index, part_index), [0, 0])[ + 0 if boundary == "start" else 1 + ] = target_cursor + source_cursor = match.end() + unmarked_parts.append(marked_text[source_cursor:]) + unmarked_text = "".join(unmarked_parts) + for key, bounds in char_bounds.items(): + leading, trailing = part_whitespace[key] + if ( + leading + and unmarked_text[max(0, bounds[0] - len(leading)) : bounds[0]] + == leading + ): + bounds[0] -= len(leading) + if ( + trailing + and unmarked_text[bounds[1] : bounds[1] + len(trailing)] + == trailing + ): + bounds[1] += len(trailing) + try: + encoded = cast(_OffsetTokenizer, resolved_tokenizer)( + unmarked_text, + add_special_tokens=False, + return_offsets_mapping=True, + ) + except (TypeError, NotImplementedError): + encoded = None + encoded_data = _string_dict(encoded) + raw_offsets = ( + encoded_data.get("offset_mapping") + if encoded_data is not None + else None + ) + if ( + encoded is not None + and _ids(encoded) == rendered + and isinstance(raw_offsets, list) + and len(raw_offsets) == len(rendered) + ): + offsets: list[tuple[int, int]] = [] + for value in raw_offsets: + if ( + not isinstance(value, (list, tuple)) + or len(value) != 2 + or not all(isinstance(item, int) for item in value) + ): + break + offsets.append((value[0], value[1])) + if len(offsets) == len(rendered): + token_bounds: dict[tuple[int, int], tuple[int, int]] = {} + token_cursor = 0 + for key, (char_start, char_end) in sorted( + char_bounds.items(), key=lambda item: item[1] + ): + while ( + token_cursor < len(offsets) + and offsets[token_cursor][1] <= char_start + ): + token_cursor += 1 + token_end = token_cursor + while ( + token_end < len(offsets) + and offsets[token_end][0] < char_end + ): + token_end += 1 + if char_start == char_end: + token_bounds[key] = (token_cursor, token_cursor) + elif ( + token_end > token_cursor + and offsets[token_cursor][0] >= char_start + and offsets[token_end - 1][1] <= char_end + ): + token_bounds[key] = (token_cursor, token_end) + token_cursor = token_end + for message_index, part_count in part_counts.items(): + bounds = [ + token_bounds[(message_index, part_index)] + for part_index in range(part_count) + if (message_index, part_index) in token_bounds + ] + if len(bounds) == part_count: + marked_part_bounds[message_index] = bounds + marked_bounds[message_index] = ( + bounds[0][0], + bounds[-1][1], + ) + for message_index, bounds in list(marked_part_bounds.items()): + whitespace_parts = [ + (part_index, text) + for part_index, (_, text) in enumerate( + _chat_message_parts(messages[message_index]) + ) + if text and not text.strip() + ] + for part_index, _ in whitespace_parts: + empty_messages = deepcopy(messages) + groups = _chat_message_text_slot_groups(empty_messages[message_index]) + if part_index >= len(groups): + marked_bounds.pop(message_index, None) + marked_part_bounds.pop(message_index, None) + break + for container, key in groups[part_index]: + container[key] = "" + try: + empty_render = render( + empty_messages, + add_generation_prompt=not ends_with_assistant, + ) + except Exception: + marked_bounds.pop(message_index, None) + marked_part_bounds.pop(message_index, None) + break + if empty_render == rendered: + continue + anchor = bounds[part_index][0] + start = anchor - (len(rendered) - len(empty_render)) + span = (start, anchor) + if ( + start < 0 + or rendered[: span[0]] + rendered[span[1] :] != empty_render + ): + marked_bounds.pop(message_index, None) + marked_part_bounds.pop(message_index, None) + break + bounds[part_index] = span + marked_bounds[message_index] = (bounds[0][0], bounds[-1][1]) + + probed_bounds: dict[int, tuple[int, int]] = {} + if not direct_bounds: + for message_index, (message, source) in enumerate( + zip(messages, history.message_sources, strict=True) + ): + if ( + message_index in marked_bounds + or message.get("role") != "assistant" + or source is None + or not _source_is_sampled(source) + ): + continue + parts = _chat_message_parts(message) + if not parts: + exact_output, _ = source_output_tokens(source) + if exact_output: + try: + prefix = render( + messages[:message_index], + add_generation_prompt=True, + ) + completed = render( + messages[: message_index + 1], + add_generation_prompt=False, + ) + except Exception: + continue + if ( + completed[: len(prefix)] == prefix + and rendered[: len(completed)] == completed + ): + probed_bounds[message_index] = ( + len(prefix), + len(prefix), + ) + continue + if parts and all(part == "tool_call" for part, _ in parts): + probe_messages = deepcopy(messages) + tool_calls = probe_messages[message_index].get("tool_calls") + if isinstance(tool_calls, list): + existing_values: set[str] = set() + for existing_call in tool_calls: + if not isinstance(existing_call, dict): + continue + existing_function = existing_call.get("function") + if isinstance(existing_function, dict): + existing_values.update( + value + for value in existing_function.values() + if isinstance(value, str) + ) + for call_index, call in enumerate(tool_calls): + if not isinstance(call, dict): + continue + raw_function = call.get("function") + if not isinstance(raw_function, dict): + continue + function = cast(dict[str, Any], raw_function) + probe_name = ( + f"art_trajectory_probe_{id(probe_messages):x}_{call_index}" + ) + probe_arguments = json.dumps({probe_name: True}) + while ( + probe_name in existing_values + or probe_arguments in existing_values + ): + probe_name += "_" + probe_arguments = json.dumps({probe_name: True}) + function["name"] = probe_name + function["arguments"] = probe_arguments + try: + probe = render( + probe_messages, + add_generation_prompt=not ends_with_assistant, + ) + except Exception: + probe = [] + if span := differing_span(probe): + probed_bounds[message_index] = span + continue + if len(parts) == 1: + try: + prefix = render( + messages[:message_index], add_generation_prompt=True + ) + except Exception: + prefix = [] + local = part_ids(parts[0][1]) + if rendered == [*prefix, *local]: + probed_bounds[message_index] = ( + len(prefix), + len(rendered), + ) + continue + try: + completed = render( + messages[: message_index + 1], + add_generation_prompt=False, + ) + except Exception: + completed = [] + if ( + completed == [*prefix, *local] + and rendered[: len(completed)] == completed + ): + probed_bounds[message_index] = ( + len(prefix), + len(completed), + ) + continue + probe_messages = deepcopy(messages) + slot_groups = _chat_message_text_slot_groups(probe_messages[message_index]) + if len(slot_groups) != 1 or len(slot_groups[0]) != 1: + continue + container, key = slot_groups[0][0] + original = str(container[key]) + leading = original[: len(original) - len(original.lstrip())] + trailing = original[len(original.rstrip()) :] + container[key] = ( + leading + f"ART_TRAJECTORY_{id(probe_messages):x}_PROBE" + trailing + ) + try: + probe = render( + probe_messages, + add_generation_prompt=not ends_with_assistant, + ) + except Exception: + continue + if span := differing_span(probe): + probed_bounds[message_index] = span + + sampled_message_count = sum( + message.get("role") == "assistant" + and source is not None + and _source_is_sampled(source) + for message, source in zip(messages, history.message_sources, strict=True) + ) + for message_index, (message, source) in enumerate( + zip(messages, history.message_sources, strict=True) + ): + parts = _chat_message_parts(message) + sampled = ( + message.get("role") == "assistant" + and source is not None + and _source_is_sampled(source) + ) + full_exact, full_logprobs = ( + source_output_tokens(source) if sampled else (None, []) + ) + if sampled and not parts and not full_exact: + continue + complete_sampled_message = ( + sampled + and source is not None + and _source_covers_complete_sampled_message(message, source) + ) + source_boundary = False + generation_start: int | None = None + sampled_bounds: tuple[int, int] | None = None + content_bounds_proven = False + if sampled: + if direct_bounds: + sampled_bounds = direct_bounds[message_index] + content_bounds_proven = True + elif message_index in marked_bounds: + sampled_bounds = marked_bounds[message_index] + content_bounds_proven = True + elif message_index in probed_bounds: + sampled_bounds = probed_bounds[message_index] + content_bounds_proven = True + else: + assert source is not None + source_prompt = source_prompt_tokens(source) + source_context_matches = source_matches_context(source) + if ( + source_context_matches + and source_prompt + and rendered[: len(source_prompt)] == source_prompt + ): + search_cursor = max(search_cursor, len(source_prompt)) + source_boundary = True + generation_start = len(source_prompt) + sampled_bounds = (generation_start, len(rendered)) + elif sampled_message_count == 1: + prompt_render = render( + messages[:message_index], add_generation_prompt=True + ) + if rendered[: len(prompt_render)] != prompt_render: + raise ValueError( + "Could not locate a sampled history message in the " + "rendered history" + ) + generation_start = len(prompt_render) + sampled_bounds = (generation_start, len(rendered)) + else: + raise ValueError( + "Could not prove a sampled history message boundary with this " + "tokenizer" + ) + full_matches = ( + locations(full_exact, search_cursor) if sampled and full_exact else [] + ) + first_part_matches = ( + locations(part_ids(parts[0][1]), search_cursor) if parts else [] + ) + if sampled_bounds is not None: + lower, upper = sampled_bounds + full_matches = [ + match + for match in full_matches + if match[0] >= lower and match[1] <= upper + ] + first_part_matches = [ + match + for match in first_part_matches + if match[0] >= lower and match[1] <= upper + ] + if ( + full_exact is not None + and full_matches + and ( + (source_boundary and full_matches[0][0] == search_cursor) + or ( + complete_sampled_message + and len(full_matches) == 1 + and first_part_matches + and full_matches[0][0] == first_part_matches[0][0] + ) + ) + ): + span = full_matches[0] + start, end = span + replacements.append( + ( + start, + end, + full_exact, + full_logprobs + if len(full_logprobs) == len(full_exact) + else [math.nan] * len(full_exact), + True, + _sampled_source_key(source), + source, + ) + ) + search_cursor = end + continue + + source_exchange = getattr(source, "exchange", None) + multi_generation_response = ( + isinstance(source_exchange, ResponsesExchange) + and len(_response_generations(source_exchange.response)) > 1 + ) + if sampled and not content_bounds_proven: + raise ValueError( + "Could not uniquely locate or prove the sampled content boundary " + "with this tokenizer" + ) + if ( + sampled + and full_exact is None + and message_index in probed_bounds + and parts + and all(part == "tool_call" for part, _ in parts) + ): + assert source is not None and sampled_bounds is not None + start, end = sampled_bounds + replacements.append( + ( + start, + end, + rendered[start:end], + [math.nan] * (end - start), + False, + _sampled_source_key(source), + source, + ) + ) + search_cursor = end + continue + if ( + complete_sampled_message + and full_exact is not None + and ( + multi_generation_response + or len(parts) != 1 + or len(full_exact) != len(part_ids(parts[0][1])) + ) + ): + if not parts and sampled_bounds is not None: + start = end = sampled_bounds[0] + elif sampled_bounds is None: + raise ValueError( + "Could not locate a complete sampled message in the rendered history" + ) + elif sampled_bounds[0] == sampled_bounds[1]: + if not content_bounds_proven: + raise ValueError( + "Could not locate a complete sampled message in the rendered " + "history" + ) + start = end = sampled_bounds[0] + else: + start, end = sampled_bounds + if parts and len(parts) == 1 and parts[0][0] == "content": + visible_matches = [ + match + for match in locations(part_ids(parts[0][1]), start) + if match[1] <= end + ] + if len(visible_matches) != 1 or ( + not content_bounds_proven + and generation_start is not None + and visible_matches[0][0] != generation_start + ): + raise ValueError( + "Could not prove the sampled content boundary in the " + "rendered history" + ) + start, end = visible_matches[0] + elif generation_start is not None and ( + multi_generation_response or len(parts) != 1 or parts[0][0] != "content" + ): + start = generation_start + replacements.append( + ( + start, + end, + full_exact, + full_logprobs + if len(full_logprobs) == len(full_exact) + else [math.nan] * len(full_exact), + True, + _sampled_source_key(source), + source, + ) + ) + if ( + message_index < len(messages) - 1 + and isinstance(source_exchange, ChatCompletionsExchange) + and rendered[start:end] != full_exact + ): + _warn_prefix_retokenization() + search_cursor = end + continue + + replacement_start = len(replacements) + for part_index, (part, text) in enumerate(parts): + if not sampled and text not in sampled_texts: + continue + local = part_ids(text) + if not local: + continue + proven_part_bounds = marked_part_bounds.get(message_index) + span = ( + proven_part_bounds[part_index] + if proven_part_bounds is not None + else next(iter(locations(local, search_cursor)), None) + ) + if span is None: + if not sampled: + continue + raise ValueError( + "Could not locate a history message in the rendered history" + ) + if sampled_bounds is not None: + lower, upper = sampled_bounds + if proven_part_bounds is None: + bounded_matches = [ + match + for match in locations(local, max(lower, search_cursor)) + if match[1] <= upper + ] + if len(bounded_matches) != 1: + raise ValueError( + "Could not uniquely locate a sampled history message in " + "the rendered history" + ) + span = bounded_matches[0] + start, end = span + search_cursor = end + if not sampled: + continue + assert source is not None + exact, logprobs = _chat_source_tokens( + source, + text, + part=part, + full_tokens=(full_exact, full_logprobs), + ) + if exact is not None and rendered[start : start + len(exact)] == exact: + end = start + len(exact) + search_cursor = end + replacement = exact if exact is not None else rendered[start:end] + if exact is None and not logprobs: + exchange = getattr(source, "exchange", None) + if isinstance( + exchange, + (ChatCompletionsExchange, ResponsesExchange, MessagesExchange), + ): + evidence = _visible_token_evidence( + tokenizer, + exchange, + source=source, + sampled_text=text, + ) + if evidence is not None: + replacement, logprobs = evidence + else: + logprobs = ( + _align_visible_logprobs( + tokenizer, + replacement, + exchange, + source=source, + sampled_text=text, + ) + or [] + ) + replacements.append( + ( + start, + end, + replacement, + logprobs + if len(logprobs) == len(replacement) + else [math.nan] * len(replacement), + exact is not None, + _sampled_source_key(source), + source, + ) + ) + message_replacements = replacements[replacement_start:] + if ( + sampled + and message_replacements + and all(part == "tool_call" for part, _ in parts) + and not ( + all(replacement[4] for replacement in message_replacements) + and all( + left[1] == right[0] + for left, right in zip( + message_replacements, + message_replacements[1:], + strict=False, + ) + ) + ) + ): + start = message_replacements[0][0] + end = message_replacements[-1][1] + del replacements[replacement_start:] + replacements.append( + ( + start, + end, + rendered[start:end], + [math.nan] * (end - start), + False, + message_replacements[0][5], + message_replacements[0][6], + ) + ) + if sampled and not parts and full_exact is not None: + raise ValueError( + "Could not locate exact sampled output in the rendered history" + ) + + token_ids: list[int] = [] + logprobs: list[float] = [] + flags: list[TokenFlag] = [] + source_keys: list[_SampledSourceKey | None] = [] + sources: dict[_SampledSourceKey, object] = {} + cursor = 0 + for ( + start, + end, + replacement, + replacement_logprobs, + exact, + source_key, + source, + ) in sorted(replacements, key=lambda item: (item[0], item[1])): + if start < cursor: + raise ValueError("Rendered assistant source spans overlap") + token_ids.extend(rendered[cursor:start]) + logprobs.extend([math.nan] * (start - cursor)) + flags.extend([TokenFlag(0)] * (start - cursor)) + source_keys.extend([None] * (start - cursor)) + token_ids.extend(replacement) + logprobs.extend(replacement_logprobs) + flag = TokenFlag.SAMPLED | (TokenFlag.EXACT if exact else TokenFlag(0)) + flags.extend([flag] * len(replacement)) + source_keys.extend([source_key] * len(replacement)) + sources[source_key] = source + cursor = end + token_ids.extend(rendered[cursor:]) + logprobs.extend([math.nan] * (len(rendered) - cursor)) + flags.extend([TokenFlag(0)] * (len(rendered) - cursor)) + source_keys.extend([None] * (len(rendered) - cursor)) + for index in range(min(exact_prefix_length, len(flags))): + flags[index] |= TokenFlag.EXACT + if history.model is None: + raise ValueError("History tokenization requires a model") + tokenized = TokenizedHistory( + model=history.model, + token_ids=token_ids, + logprobs=logprobs, + flags=flags, + ) + if _trace is not None: + _trace.set(tokenized, source_keys, sources) + return tokenized + + +def _tokenize_completions_token_history( + history: CompletionsTokenHistory, + *, + _trace: _TraceBuilder | None = None, +) -> TokenizedHistory: + if any( + span.start < 0 or span.end <= span.start or span.end > len(history.prompt) + for span in history.prompt_sources + ): + raise ValueError("Completions token source spans are out of bounds") + if not _spans_are_exhaustive(len(history.prompt), history.prompt_sources): + raise ValueError( + "Completions token source spans must exhaustively cover prompt" + ) + _validate_completions_sources( + model=history.model, + source_spans=history.prompt_sources, + sampled_spans=history.sampled_spans, + ) + + flags = [TokenFlag(0)] * len(history.prompt) + logprobs = [math.nan] * len(history.prompt) + source_keys: list[_SampledSourceKey | None] = [None] * len(history.prompt) + sources: dict[_SampledSourceKey, object] = {} + for start, end in history.sampled_spans: + if start < 0 or end <= start or end > len(history.prompt): + raise ValueError("Completions sampled spans are out of bounds") + flags[start:end] = [TokenFlag.SAMPLED] * (end - start) + for span in history.prompt_sources: + if span.source is None: + continue + if span.source.choice_index is not None: + source_key = _sampled_source_key(span.source) + source_keys[span.start : span.end] = [source_key] * (span.end - span.start) + sources[source_key] = span.source + flags[span.start : span.end] = [ + flag | TokenFlag.EXACT for flag in flags[span.start : span.end] + ] + prompt, completion, prompt_logprobs, completion_logprobs = ( + _completion_source_evidence(span.source) + ) + selected = prompt if span.source.choice_index is None else completion + selected_logprobs = ( + prompt_logprobs if span.source.choice_index is None else completion_logprobs + ) + if selected != history.prompt[span.start : span.end]: + if span.source.choice_index is not None: + raise ValueError( + "Completions sampled output no longer matches its source exchange" + ) + flags[span.start : span.end] = [ + flag & ~TokenFlag.EXACT for flag in flags[span.start : span.end] + ] + continue + if len(selected_logprobs) == span.end - span.start: + logprobs[span.start : span.end] = selected_logprobs + tokenized = TokenizedHistory( + model=history.model, + token_ids=list(history.prompt), + logprobs=logprobs, + flags=flags, + ) + if _trace is not None: + _trace.set(tokenized, source_keys, sources) + return tokenized + + +def _completion_visible_logprobs( + source: CompletionsSource, + text: str, + tokenizer: Tokenizer, + token_ids: list[int], +) -> list[float] | None: + choice = next( + item + for item in source.exchange.response.choices + if item.index == source.choice_index + ) + data = _dump(choice.logprobs) + raw_tokens = data.get("tokens") + raw_logprobs = data.get("token_logprobs") + if not isinstance(raw_tokens, list) or not isinstance(raw_logprobs, list): + return None + values = list(zip(raw_tokens, raw_logprobs, strict=False)) + from ._history import _completion_prompts + + request_prompts = _completion_prompts(source.exchange.request.get("prompt")) + request_prompt = request_prompts[source.prompt_index] + if source.exchange.request.get("echo") is True and isinstance(request_prompt, str): + consumed = "" + cursor = 0 + while cursor < len(values) and len(consumed) < len(request_prompt): + token_text = values[cursor][0] + if not isinstance(token_text, str): + return None + consumed += token_text + cursor += 1 + if consumed != request_prompt: + return None + values = values[cursor:] + if ( + "".join(token_text for token_text, _ in values if isinstance(token_text, str)) + != text + ): + return None + aligned_ids: list[int] = [] + aligned_logprobs: list[float] = [] + for token_text, logprob in values: + if ( + not isinstance(token_text, str) + or not isinstance(logprob, (int, float)) + or isinstance(logprob, bool) + ): + return None + encoded = _ids(tokenizer(token_text, add_special_tokens=False)) + if len(encoded) != 1: + return None + aligned_ids.append(encoded[0]) + aligned_logprobs.append(float(logprob)) + return aligned_logprobs if aligned_ids == token_ids else None + + +def _tokenize_completions_string_history( + history: CompletionsStringHistory, + *, + base_model: str | None, + tokenizer: Tokenizer | None, + _trace: _TraceBuilder | None = None, +) -> TokenizedHistory: + if not _spans_are_exhaustive(len(history.prompt), history.prompt_sources): + raise ValueError( + "Completions string source spans must exhaustively cover prompt" + ) + _validate_completions_sources( + model=history.model, + source_spans=history.prompt_sources, + sampled_spans=history.sampled_spans, + ) + sampled = [False] * len(history.prompt) + for start, end in history.sampled_spans: + if start < 0 or end <= start or end > len(history.prompt): + raise ValueError("Completions sampled spans are out of bounds") + sampled[start:end] = [True] * (end - start) + + config: _TokenizerConfig | None = None + + def resolved_tokenizer() -> Tokenizer: + nonlocal config, tokenizer + if tokenizer is None: + config = config or _tokenizer_config(history.model, base_model) + tokenizer = _load_tokenizer(config) + return tokenizer + + token_ids: list[int] = [] + logprobs: list[float] = [] + flags: list[TokenFlag] = [] + source_keys: list[_SampledSourceKey | None] = [] + sources: dict[_SampledSourceKey, object] = {} + for span in history.prompt_sources: + text = history.prompt[span.start : span.end] + source = span.source + exact: list[int] | None = None + source_logprobs: list[float] = [] + is_sampled = any(sampled[span.start : span.end]) + if source is not None: + prompt, completion, prompt_logprobs, completion_logprobs = ( + _completion_source_evidence(source) + ) + if source.choice_index is None: + from ._history import _completion_prompts + + request_prompts = _completion_prompts( + source.exchange.request.get("prompt") + ) + original = request_prompts[source.prompt_index] + if not isinstance(original, str): + raise ValueError( + "A string Completions history cannot reference a token prompt" + ) + if text == original: + exact = prompt + source_logprobs = prompt_logprobs + else: + choice = next( + item + for item in source.exchange.response.choices + if item.index == source.choice_index + ) + expected = choice.text + from ._history import _completion_prompts + + request_prompts = _completion_prompts( + source.exchange.request.get("prompt") + ) + request_prompt = request_prompts[source.prompt_index] + if source.exchange.request.get("echo") is True: + if not isinstance(request_prompt, str) or not expected.startswith( + request_prompt + ): + raise ValueError( + "Cannot locate echoed Completions prompt boundary" + ) + expected = expected[len(request_prompt) :] + if text != expected: + raise ValueError( + "Completions history text no longer matches its source exchange" + ) + exact = completion + source_logprobs = completion_logprobs + ids = ( + exact + if exact is not None + else _ids(resolved_tokenizer()(text, add_special_tokens=False)) + ) + token_ids.extend(ids) + if exact is not None and len(source_logprobs) == len(ids): + logprobs.extend(source_logprobs) + else: + visible = ( + _completion_visible_logprobs(source, text, resolved_tokenizer(), ids) + if source is not None and source.choice_index is not None + else None + ) + logprobs.extend(visible or [math.nan] * len(ids)) + flag = TokenFlag.SAMPLED if is_sampled else TokenFlag(0) + if exact is not None: + flag |= TokenFlag.EXACT + flags.extend([flag] * len(ids)) + if is_sampled: + if source is None: + raise AssertionError("Sampled Completions span has no source") + source_key = _sampled_source_key(source) + source_keys.extend([source_key] * len(ids)) + sources[source_key] = source + else: + source_keys.extend([None] * len(ids)) + tokenized = TokenizedHistory( + model=history.model, + token_ids=token_ids, + logprobs=logprobs, + flags=flags, + ) + if _trace is not None: + _trace.set(tokenized, source_keys, sources) + return tokenized + + +def _completion_source_evidence( + source: CompletionsSource, +) -> tuple[list[int] | None, list[int] | None, list[float], list[float]]: + from ._history import _completion_choice_groups + + prompt_groups = _completion_choice_groups(source.exchange) + prompt_index = _source_index( + source.prompt_index, + length=len(prompt_groups), + field="Completions prompt source index", + ) + if source.choice_index is None: + selected = prompt_groups[prompt_index][0] + else: + if ( + not isinstance(source.choice_index, int) + or isinstance(source.choice_index, bool) + or source.choice_index < 0 + ): + raise ValueError("Completions choice source index is invalid") + selected = next( + ( + choice + for choice in prompt_groups[prompt_index] + if choice.index == source.choice_index + ), + None, + ) + if selected is None: + raise ValueError("Completions choice source does not belong to its prompt") + return _completion_evidence( + source.exchange.response.model_copy(update={"choices": [selected]}), + echo=source.exchange.request.get("echo") is True, + empty_prompt_is_exact=source.exchange.request.get("prompt") in ("", []), + ) + + +def _validate_completions_sources( + *, + model: str, + source_spans: Sequence[CompletionsTokenSourceSpan | CompletionsStringSourceSpan], + sampled_spans: Sequence[tuple[int, int]], +) -> None: + expected_sampled: list[tuple[int, int]] = [] + for span in source_spans: + source = getattr(span, "source", None) + if source is None: + continue + if not isinstance(source, CompletionsSource): + raise ValueError("Completions history has an invalid source") + if source.exchange.model != model: + raise ValueError( + "Completions history model no longer matches its source exchange" + ) + if source.choice_index is not None: + expected_sampled.append((span.start, span.end)) + if list(sampled_spans) != expected_sampled: + raise ValueError( + "Completions sampled spans must exactly match choice-backed source spans" + ) + + +def _spans_are_exhaustive(length: int, spans: Sequence[object]) -> bool: + cursor = 0 + for span in spans: + start = getattr(span, "start", None) + end = getattr(span, "end", None) + if ( + not isinstance(start, int) + or not isinstance(end, int) + or start != cursor + or end <= start + or end > length + ): + return False + cursor = end + return cursor == length + + +def tokenize_history( + history: History | LegacyHistory, + *, + model: str | None, + base_model: str | None, + tokenizer: Tokenizer | None, + chat_template: str | None, + chat_template_kwargs: Mapping[str, object] | None, + _trace: _TraceBuilder | None = None, + _projection_validated: bool = False, +) -> TokenizedHistory: + if isinstance(history, LegacyHistory): + if model is None: + raise ValueError("Legacy history tokenization requires model=") + return _legacy_tokenize(history, model=model) + if model is None: + raise ValueError("History tokenization requires a model") + if isinstance(history, CompletionsTokenHistory): + return _tokenize_completions_token_history(history, _trace=_trace) + if isinstance(history, CompletionsStringHistory): + return _tokenize_completions_string_history( + history, + base_model=base_model, + tokenizer=tokenizer, + _trace=_trace, + ) + _validate_history_sources(history) + override_requires_render = ( + chat_template is not None + and chat_template != getattr(history, "chat_template", None) + ) or ( + chat_template_kwargs is not None + and dict(chat_template_kwargs) + != (getattr(history, "chat_template_kwargs", None) or {}) + ) + render_state = ( + _HistoryRenderState(needs_render=False, projection_matches=True) + if _projection_validated + else _history_render_state(history) + ) + needs_render = render_state.needs_render or override_requires_render + if isinstance(history, ResponsesHistory) and not needs_render: + if exact := _tokenize_exact_responses_history( + history, base_model=base_model, tokenizer=tokenizer, _trace=_trace + ): + return exact + if isinstance(history, ChatCompletionsHistory): + if ( + not override_requires_render + and not render_state.context_changed + and render_state.projection_matches is not False + and ( + exact := _tokenize_exact_projected_chat_history( + history, + projection_validated=( + _projection_validated or render_state.projection_matches is True + ), + _trace=_trace, + ) + ) + ): + return exact + return _tokenize_chat_view( + history, + base_model=base_model, + tokenizer=tokenizer, + chat_template=chat_template, + chat_template_kwargs=chat_template_kwargs, + _projection_matches=( + True if _projection_validated else render_state.projection_matches + ), + _trace=_trace, + ) + if isinstance(history, AnthropicMessagesHistory) and needs_render: + converted = history.as_chat_completions_history() + if ( + not override_requires_render + and not render_state.context_changed + and ( + render_state.projection_matches + if render_state.projection_matches is not None + else _history_matches_projection(history) + ) + and ( + exact := _tokenize_exact_projected_chat_history( + converted, + projection_validated=True, + _trace=_trace, + ) + ) + ): + return exact + return _tokenize_chat_view( + converted, + base_model=base_model, + tokenizer=tokenizer, + chat_template=chat_template, + chat_template_kwargs=chat_template_kwargs, + _trace=_trace, + ) + if isinstance(history, ResponsesHistory) and needs_render: + return _tokenize_chat_view( + history.as_chat_completions_history(), + base_model=base_model, + tokenizer=tokenizer, + chat_template=chat_template, + chat_template_kwargs=chat_template_kwargs, + _trace=_trace, + ) + trajectory = _trajectory_from_history(history) + history_template = getattr(history, "chat_template", None) + history_kwargs = getattr(history, "chat_template_kwargs", None) + return _tokenize_exchange_trajectory( + trajectory, + base_model, + model=model, + chat_template=( + chat_template if chat_template is not None else history_template + ), + chat_template_kwargs={ + **(history_kwargs or {}), + **(chat_template_kwargs or {}), + } + or None, + tokenizer_instance=tokenizer, + _trace=_trace, + ) + + +def _materialize_trajectory( + tokenized: TokenizedHistory, trajectory: Trajectory +) -> TokenizedTrajectory: + return TokenizedTrajectory( + **tokenized.model_dump(), + reward=trajectory.reward, + metrics=dict(trajectory.metrics), + metadata=dict(trajectory.metadata), + ) + + +def tokenize_trajectory( + trajectory: Trajectory, + *, + multi_history: bool, + model: str | None, + base_model: str | None, + tokenizer: Tokenizer | None, + chat_template: str | None, + chat_template_kwargs: Mapping[str, object] | None, +) -> TokenizedTrajectory | TokenizedMultiHistoryTrajectory: + histories = trajectory.histories(model=model) + if not multi_history: + if len(histories) != 1: + selected_models = { + history.model + for history in histories + if not isinstance(history, LegacyHistory) + } + if model is None and len(selected_models) > 1: + raise ValueError( + "Trajectory tokenization requires exactly one model; pass model= to select one" + ) + raise ValueError( + f"Trajectory tokenization requires exactly one history; found {len(histories)}" + ) + tokenized = [ + tokenize_history( + history, + model=model if isinstance(history, LegacyHistory) else history.model, + base_model=base_model, + tokenizer=tokenizer, + chat_template=chat_template, + chat_template_kwargs=chat_template_kwargs, + _projection_validated=not isinstance(history, LegacyHistory), + ) + for history in histories + ] + if not multi_history: + return _materialize_trajectory(tokenized[0], trajectory) + return TokenizedMultiHistoryTrajectory( + histories=tokenized, + reward=trajectory.reward, + metrics=dict(trajectory.metrics), + metadata=dict(trajectory.metadata), + ) + + +def _tokenize_trajectory_with_trace( + trajectory: Trajectory, + *, + model: str | None = None, + base_model: str | None = None, + tokenizer: Tokenizer | None = None, + chat_template: str | None = None, + chat_template_kwargs: Mapping[str, object] | None = None, +) -> tuple[ + TokenizedMultiHistoryTrajectory, + list[_HistoryTokenizationTrace], +]: + if not trajectory.exchanges: + raise ValueError("Private exchange tokenization trace requires exchanges") + histories = trajectory.histories(model=model) + tokenized_histories: list[TokenizedHistory] = [] + traces: list[_HistoryTokenizationTrace] = [] + for history in histories: + if isinstance(history, LegacyHistory): + raise AssertionError( + "Exchange trajectories cannot produce legacy histories" + ) + trace_builder = _TraceBuilder() + tokenized = tokenize_history( + history, + model=history.model, + base_model=base_model, + tokenizer=tokenizer, + chat_template=chat_template, + chat_template_kwargs=chat_template_kwargs, + _trace=trace_builder, + _projection_validated=True, + ) + if trace_builder.trace is None: + raise AssertionError("Exchange tokenization did not produce a source trace") + tokenized_histories.append(tokenized) + traces.append(trace_builder.trace) + return ( + TokenizedMultiHistoryTrajectory( + histories=tokenized_histories, + reward=trajectory.reward, + metrics=dict(trajectory.metrics), + metadata=dict(trajectory.metadata), + ), + traces, + ) + + +def tokenize_group( + group: TrajectoryGroup, + *, + multi_history: bool, + model: str | None, + base_model: str | None, + tokenizer: Tokenizer | None, + chat_template: str | None, + chat_template_kwargs: Mapping[str, object] | None, +) -> ( + TokenizedTrajectoryGroup[TokenizedTrajectory] + | TokenizedTrajectoryGroup[TokenizedMultiHistoryTrajectory] +): + trajectories = [ + tokenize_trajectory( + trajectory, + multi_history=multi_history, + model=model, + base_model=base_model, + tokenizer=tokenizer, + chat_template=chat_template, + chat_template_kwargs=chat_template_kwargs, + ) + for trajectory in group.trajectories + ] + if multi_history: + return TokenizedTrajectoryGroup[TokenizedMultiHistoryTrajectory]( + trajectories=trajectories, + metrics=dict(group.metrics), + metadata=dict(group.metadata), + ) + return TokenizedTrajectoryGroup[TokenizedTrajectory]( + trajectories=trajectories, + metrics=dict(group.metrics), + metadata=dict(group.metadata), ) diff --git a/src/art/utils/trajectory_logging.py b/src/art/utils/trajectory_logging.py index 325f576d6..eb697229a 100644 --- a/src/art/utils/trajectory_logging.py +++ b/src/art/utils/trajectory_logging.py @@ -9,7 +9,7 @@ import pydantic from art.openai import ART_MOE_ROUTING_METADATA_KEY -from art.trajectories import History, Trajectory, TrajectoryGroup +from art.trajectories import LegacyHistory, Trajectory, TrajectoryGroup if TYPE_CHECKING: import pyarrow as pa @@ -88,7 +88,7 @@ def _choice_data(item: object) -> dict[str, Any] | None: ) -def _history_data(history: History) -> dict[str, Any]: +def _history_data(history: LegacyHistory) -> dict[str, Any]: data = history.model_dump( mode="json", exclude={"messages_and_choices"}, warnings="error" ) @@ -123,7 +123,7 @@ def _trajectory_data(trajectory: Trajectory) -> dict[str, Any]: return data -def _restore_history(data: object) -> History: +def _restore_history(data: object) -> LegacyHistory: if not isinstance(data, dict): raise ValueError("Parquet additional history must be a JSON object") restored = dict(data) @@ -136,7 +136,7 @@ def _restore_history(data: object) -> History: else item for item in messages ] - return History.model_validate(restored) + return LegacyHistory.model_validate(restored) def _restore_trajectory(payload: object) -> Trajectory: diff --git a/src/art/utils/trajectory_migration.py b/src/art/utils/trajectory_migration.py index 7a2121893..10c40b1b8 100644 --- a/src/art/utils/trajectory_migration.py +++ b/src/art/utils/trajectory_migration.py @@ -19,7 +19,7 @@ import pydantic import yaml -from art.trajectories import History, Trajectory, TrajectoryGroup +from art.trajectories import LegacyHistory, Trajectory, TrajectoryGroup from art.types import Choice, Message, MessageOrChoice from art.utils.trajectory_logging import write_trajectory_groups_parquet @@ -49,7 +49,7 @@ def trajectory_group_to_dict(trajectory_group: TrajectoryGroup) -> dict[str, Any } -def history_to_dict(history: History) -> dict[str, Any]: +def history_to_dict(history: LegacyHistory) -> dict[str, Any]: messages_and_choices = [ message_or_choice_to_dict(message_or_choice) for message_or_choice in history.messages_and_choices @@ -146,8 +146,8 @@ def dict_to_trajectory(d: dict[str, Any]) -> Trajectory: ) -def dict_to_history(d: dict[str, Any]) -> History: - return History.model_validate( +def dict_to_history(d: dict[str, Any]) -> LegacyHistory: + return LegacyHistory.model_validate( { **d, "messages_and_choices": [ diff --git a/tests/integration/megatron/model_support/chat_template_rollout.py b/tests/integration/megatron/model_support/chat_template_rollout.py index 30a42148e..3ea2d897f 100644 --- a/tests/integration/megatron/model_support/chat_template_rollout.py +++ b/tests/integration/megatron/model_support/chat_template_rollout.py @@ -16,7 +16,7 @@ tokenize_trajectory, tokenize_trajectory_groups, ) -from art.trajectories import History +from art.trajectories import LegacyHistory from tests.support.chat_template_conformance_cases import ( build_chat_template_conformance_inputs, ) @@ -33,8 +33,8 @@ def _artifact_dir(base_model: str) -> Path: return path -def _history(trajectory: art.Trajectory) -> History: - return History( +def _history(trajectory: art.Trajectory) -> LegacyHistory: + return LegacyHistory( messages_and_choices=trajectory.messages_and_choices, tools=trajectory.tools, ) diff --git a/tests/integration/megatron/runtime_isolation/test_runtime_project_isolation.py b/tests/integration/megatron/runtime_isolation/test_runtime_project_isolation.py index eb7501217..4eddbb176 100644 --- a/tests/integration/megatron/runtime_isolation/test_runtime_project_isolation.py +++ b/tests/integration/megatron/runtime_isolation/test_runtime_project_isolation.py @@ -53,29 +53,65 @@ def test_runtime_server_source_contains_only_required_custom_routes() -> None: assert route in source -def test_runtime_patch_always_returns_token_ids( +def test_runtime_patch_defaults_evidence_on_and_honors_opt_out( artifact_dir: Path, ) -> None: payload = _runtime_python( "import json; " - "from art_vllm_runtime.patches import apply_vllm_runtime_patches; " - "apply_vllm_runtime_patches(); " + "from art_vllm_runtime.patches import subclass_chat_completion_request; " + "subclass_chat_completion_request(); " "from vllm.entrypoints.openai.chat_completion import protocol; " - "request = protocol.ChatCompletionRequest(" + "default_request = protocol.ChatCompletionRequest(" "model='m', messages=[{'role': 'user', 'content': 'x'}]" "); " + "explicit_false = protocol.ChatCompletionRequest(" + "model='m', messages=[{'role': 'user', 'content': 'x'}], " + "logprobs=False, top_logprobs=None, return_token_ids=False" + "); " + "explicit_none = protocol.ChatCompletionRequest(" + "model='m', messages=[{'role': 'user', 'content': 'x'}], " + "return_token_ids=None" + "); " "print(json.dumps({" - "'logprobs': request.logprobs, " - "'top_logprobs': request.top_logprobs, " - "'return_token_ids': request.return_token_ids" + "'default': {" + "'logprobs': default_request.logprobs, " + "'top_logprobs': default_request.top_logprobs, " + "'return_token_ids': default_request.return_token_ids" + "}, " + "'default_fields_set': sorted(default_request.model_fields_set), " + "'explicit_false': {" + "'logprobs': explicit_false.logprobs, " + "'top_logprobs': explicit_false.top_logprobs, " + "'return_token_ids': explicit_false.return_token_ids" + "}, " + "'explicit_false_fields_set': sorted(explicit_false.model_fields_set), " + "'explicit_none_return_token_ids': explicit_none.return_token_ids, " + "'explicit_none_fields_set': sorted(explicit_none.model_fields_set)" "}))", artifact_dir, "route_token_ids", ) - assert json.loads(payload) == { - "logprobs": True, - "top_logprobs": 0, - "return_token_ids": True, + assert json.loads(payload.splitlines()[-1]) == { + "default": { + "logprobs": True, + "top_logprobs": 0, + "return_token_ids": True, + }, + "default_fields_set": ["messages", "model"], + "explicit_false": { + "logprobs": False, + "top_logprobs": None, + "return_token_ids": False, + }, + "explicit_false_fields_set": [ + "logprobs", + "messages", + "model", + "return_token_ids", + "top_logprobs", + ], + "explicit_none_return_token_ids": None, + "explicit_none_fields_set": ["messages", "model", "return_token_ids"], } diff --git a/tests/support/chat_template_conformance_cases.py b/tests/support/chat_template_conformance_cases.py index 960bc6599..a3b15d03e 100644 --- a/tests/support/chat_template_conformance_cases.py +++ b/tests/support/chat_template_conformance_cases.py @@ -11,7 +11,7 @@ _apply_chat_template_token_ids, _messages_for_chat_template, ) -from art.trajectories import History, Trajectory, TrajectoryGroup +from art.trajectories import LegacyHistory, Trajectory, TrajectoryGroup from art.types import MessagesAndChoices, Tools @@ -133,7 +133,7 @@ def _rendered_ids( def _attach_token_metadata_to_history( tokenizer: PreTrainedTokenizerBase, - history: Trajectory | History, + history: Trajectory | LegacyHistory, ) -> None: items = history.messages_and_choices for index, item in enumerate(items): @@ -291,7 +291,7 @@ def build_chat_template_conformance_inputs( _choice_for_text("maybe", maybe_ids), ), additional_histories=[ - History( + LegacyHistory( messages_and_choices=_messages_and_choices( {"role": "user", "content": "Previous turn."}, _choice_for_text("prior yes", prior_yes_ids), @@ -306,7 +306,7 @@ def build_chat_template_conformance_inputs( _choice_for_text("yes", yes_ids), ), additional_histories=[ - History( + LegacyHistory( messages_and_choices=_messages_and_choices( {"role": "user", "content": "Previous turn."}, _choice_for_text("prior yes", prior_yes_ids), diff --git a/tests/unit/test_exchange_training_model_selection.py b/tests/unit/test_exchange_training_model_selection.py new file mode 100644 index 000000000..f402407bf --- /dev/null +++ b/tests/unit/test_exchange_training_model_selection.py @@ -0,0 +1,996 @@ +from __future__ import annotations + +from collections.abc import Mapping +from datetime import datetime, timedelta +from pathlib import Path +from types import SimpleNamespace +from typing import SupportsIndex, cast, overload +from unittest.mock import patch + +import numpy as np +from openai.types import Completion +from openai.types.chat import ChatCompletion, ChatCompletionMessageParam +import pytest +from transformers import PreTrainedTokenizerBase + +import art +from art import TrainableModel +from art.dev.model import InternalModelConfig +from art.local import LocalBackend +from art.openai import ART_MOE_ROUTING_METADATA_KEY +from art.preprocessing.moe_routing import MoeRouteSegments +from art.preprocessing.tokenize import ( + TokenizedResult, + _chat_choice_trace, + tokenize_trajectory_groups, +) +from art.tinker_native.data import trajectory_groups_to_datums +from art.trajectories import ( + ChatCompletionsExchange, + ChatCompletionsHistory, + ChatCompletionsMessageSource, + ChatCompletionsRequest, + CompletionsExchange, + CompletionsRequest, + TokenizedMultiHistoryTrajectory, + Tokenizer, +) +from art.trajectories import _tokenize as trajectory_tokenization +from art.trajectories._selection import ( + automatic_training_model_selector, + resolve_training_model, +) +from art.trajectories._tokenize import ( + _first_introduction_mask, + _HistoryTokenizationTrace, +) + + +def _exchange(model: str, output_token: int) -> ChatCompletionsExchange: + response = ChatCompletion.model_validate( + { + "id": f"chatcmpl-{output_token}", + "object": "chat.completion", + "created": 1, + "model": model, + "choices": [ + { + "index": 0, + "finish_reason": "stop", + "message": {"role": "assistant", "content": "answer"}, + "prompt_token_ids": [1], + "token_ids": [output_token], + "logprobs": { + "content": [ + { + "token": f"token_id:{output_token}", + "logprob": -0.1, + "bytes": [], + "top_logprobs": [], + } + ] + }, + } + ], + } + ) + start = datetime(2026, 1, 1) + return ChatCompletionsExchange( + request=ChatCompletionsRequest( + model=model, + messages=[{"role": "user", "content": "question"}], + ), + response=response, + start_time=start, + end_time=start + timedelta(seconds=1), + ) + + +def _empty_prompt_completion_exchange( + output_token_ids: list[int], +) -> CompletionsExchange: + start = datetime(2026, 1, 1) + return CompletionsExchange( + request=CompletionsRequest(model="policy", prompt=""), + response=Completion.model_validate( + { + "id": "cmpl-empty", + "object": "text_completion", + "created": 1, + "model": "policy", + "choices": [ + { + "index": 0, + "finish_reason": "stop", + "text": "answer", + "prompt_token_ids": [], + "token_ids": output_token_ids, + "logprobs": { + "tokens": [ + f"token_id:{token_id}" for token_id in output_token_ids + ], + "token_logprobs": [-0.1] * len(output_token_ids), + "top_logprobs": [{}] * len(output_token_ids), + "text_offset": list(range(len(output_token_ids))), + }, + } + ], + } + ), + start_time=start, + end_time=start + timedelta(seconds=1), + ) + + +def _routed_exchange( + *, + prompt_token_ids: list[int], + output_token: int, + messages: list[ChatCompletionMessageParam], + content: str, +) -> ChatCompletionsExchange: + exchange = _exchange("policy", output_token) + exchange.request["messages"] = messages + choice = exchange.response.choices[0] + choice.message.content = content + extra = choice.model_extra + assert extra is not None + extra["prompt_token_ids"] = prompt_token_ids + extra[ART_MOE_ROUTING_METADATA_KEY] = { + "prompt_token_ids": prompt_token_ids, + "completion_token_ids": [output_token], + "routed_experts": np.asarray( + [[[10]]] * len(prompt_token_ids) + [[[output_token * 10]]], + dtype=np.int32, + ), + } + return exchange + + +def _reasoning_stripped_group() -> art.TrajectoryGroup: + def set_choice( + exchange: ChatCompletionsExchange, + token_ids: list[int], + *, + content: str, + reasoning: str, + ) -> None: + data = exchange.response.model_dump(mode="python") + choice = data["choices"][0] + choice["message"] = { + "role": "assistant", + "content": content, + "reasoning": reasoning, + } + choice["token_ids"] = token_ids + choice["logprobs"]["content"] = [ + { + "token": f"token_id:{token_id}", + "logprob": -0.1, + "bytes": [], + "top_logprobs": [], + } + for token_id in token_ids + ] + exchange.response = ChatCompletion.model_validate(data) + extra = exchange.response.choices[0].model_extra + assert extra is not None + extra.pop(ART_MOE_ROUTING_METADATA_KEY, None) + + first = _routed_exchange( + prompt_token_ids=[1], + output_token=2, + messages=[{"role": "user", "content": "one"}], + content="first", + ) + set_choice( + first, + [2, 101, 102, 103, 104, 9], + content="first", + reasoning="long reasoning", + ) + second = _routed_exchange( + prompt_token_ids=[1, 9, 4], + output_token=5, + messages=[ + {"role": "user", "content": "one"}, + {"role": "assistant", "content": "first"}, + {"role": "user", "content": "two"}, + ], + content="second", + ) + set_choice( + second, + [5, 6], + content="second", + reasoning="short reasoning", + ) + return art.TrajectoryGroup( + [ + art.Trajectory( + exchanges=art.TrajectoryExchanges(chat_completions=[first, second]), + reward=reward, + ) + for reward in (1.0, 0.0) + ] + ) + + +def _group() -> art.TrajectoryGroup: + trajectories = [ + art.Trajectory( + exchanges=art.TrajectoryExchanges( + chat_completions=[ + _exchange("policy", 2), + _exchange("judge", 3), + ] + ), + reward=reward, + ) + for reward in (1.0, 0.0) + ] + return art.TrajectoryGroup(trajectories=trajectories) + + +def _versioned_group() -> art.TrajectoryGroup: + return art.TrajectoryGroup( + trajectories=[ + art.Trajectory( + exchanges=art.TrajectoryExchanges( + chat_completions=[ + _exchange("policy@12", 2), + _exchange("judge@4", 3), + _exchange("policy@13", 4), + ] + ), + reward=reward, + ) + for reward in (1.0, 0.0) + ] + ) + + +class _Tokenizer: + name_or_path = "base/model" + + +def test_first_introduction_mask_trains_repeated_sources_once_across_histories() -> ( + None +): + seen: set[str] = set() + assert _first_introduction_mask([None, "a", "a"], seen) == [ + False, + True, + True, + ] + assert _first_introduction_mask([None, "a", "a", None, "b"], seen) == [ + False, + False, + False, + False, + True, + ] + assert _first_introduction_mask( + [None, "a", "a", None, "b", None, "c", "c"], seen + ) == [ + False, + False, + False, + False, + False, + False, + True, + True, + ] + + +def test_preprocessing_requires_model_selection() -> None: + tokenizer = cast(PreTrainedTokenizerBase, _Tokenizer()) + with pytest.raises(ValueError, match="exactly one concrete model"): + list( + tokenize_trajectory_groups( + tokenizer, + [_group()], + allow_training_without_logprobs=False, + scale_rewards=False, + ) + ) + + results = list( + tokenize_trajectory_groups( + tokenizer, + [_group()], + allow_training_without_logprobs=False, + scale_rewards=False, + model="policy", + ) + ) + assert len(results) == 2 + assert all(result.token_ids == [1, 2] for result in results) + + +def test_tinker_requires_model_selection() -> None: + with pytest.raises(ValueError, match="exactly one concrete model"): + trajectory_groups_to_datums([_group()], renderer=None, tokenizer=None) + + datums = trajectory_groups_to_datums( + [_group()], + renderer=None, + tokenizer=None, + model="policy", + ) + assert len(datums) == 2 + assert all(datum.model_input.to_ints() == [1] for datum in datums) + + +def test_training_tokenizes_each_exchange_trajectory_once( + monkeypatch: pytest.MonkeyPatch, +) -> None: + calls = 0 + original = trajectory_tokenization._tokenize_trajectory_with_trace + + def counted( + trajectory: art.Trajectory, + *, + model: str | None = None, + base_model: str | None = None, + tokenizer: Tokenizer | None = None, + chat_template: str | None = None, + chat_template_kwargs: Mapping[str, object] | None = None, + ) -> tuple[ + TokenizedMultiHistoryTrajectory, + list[_HistoryTokenizationTrace], + ]: + nonlocal calls + calls += 1 + return original( + trajectory, + model=model, + base_model=base_model, + tokenizer=tokenizer, + chat_template=chat_template, + chat_template_kwargs=chat_template_kwargs, + ) + + monkeypatch.setattr( + trajectory_tokenization, "_tokenize_trajectory_with_trace", counted + ) + group = _group() + list( + tokenize_trajectory_groups( + cast(PreTrainedTokenizerBase, _Tokenizer()), + [group], + allow_training_without_logprobs=False, + scale_rewards=False, + model="policy", + ) + ) + assert calls == len(group.trajectories) + + calls = 0 + trajectory_groups_to_datums( + [group], + renderer=None, + tokenizer=None, + model="policy", + ) + assert calls == len(group.trajectories) + + +def test_overlength_history_does_not_claim_sources_from_fitting_history() -> None: + results = list( + tokenize_trajectory_groups( + cast(PreTrainedTokenizerBase, _Tokenizer()), + [_reasoning_stripped_group()], + allow_training_without_logprobs=False, + scale_rewards=False, + shuffle_group_trajectories=False, + drop_zero_advantage_trajectories=False, + model="policy", + _max_sequence_length=5, + ) + ) + + long = [result for result in results if len(result.token_ids) > 5] + fitting = [result for result in results if len(result.token_ids) <= 5] + assert len(long) == len(fitting) == 2 + assert all(result.assistant_mask == [0] * 7 for result in long) + assert all(result.token_ids == [1, 9, 4, 5, 6] for result in fitting) + assert all(result.assistant_mask == [0, 1, 0, 1, 1] for result in fitting) + assert all(result.weight == pytest.approx(1 / 3) for result in results) + + +def test_local_backend_trains_retained_source_after_overlength_history( + tmp_path: Path, +) -> None: + backend = LocalBackend(path=str(tmp_path)) + model = TrainableModel( + run_name="reasoning-stripped-overlength", + name="policy", + project="pipeline-tests", + base_model="test-model", + base_path=str(tmp_path), + _internal_config=InternalModelConfig(init_args={"max_seq_length": 5}), + ) + tokenizer = cast( + PreTrainedTokenizerBase, + SimpleNamespace( + name_or_path="test-model", + eos_token_id=0, + decode=lambda token_id: str(token_id), + ), + ) + backend._tokenizers[("test-model", None)] = tokenizer + backend._image_processors["test-model"] = None + + with ( + patch.object(backend, "_model_inference_name", return_value="policy"), + pytest.warns(UserWarning, match="Dropping 2 tokenized results"), + ): + packed = backend._get_packed_tensors( + model, + [_reasoning_stripped_group()], + advantage_balance=0.0, + allow_training_without_logprobs=False, + scale_rewards=False, + plot_tensors=False, + packed_sequence_length=5, + logprob_calculation_chunk_size=1, + ) + + assert packed is not None + assert packed["tokens"].tolist() == [[1, 9, 4, 5, 6]] * 2 + assert packed["assistant_mask"].tolist() == [[False, True, False, True, True]] * 2 + + +def test_training_rejects_multiple_concrete_policy_versions() -> None: + group = _versioned_group() + with pytest.raises(ValueError, match="exactly one concrete model"): + list( + tokenize_trajectory_groups( + cast(PreTrainedTokenizerBase, _Tokenizer()), + [group], + allow_training_without_logprobs=False, + scale_rewards=False, + model="policy@*", + ) + ) + + with pytest.raises(ValueError, match="exactly one concrete model"): + trajectory_groups_to_datums( + [group], + renderer=None, + tokenizer=None, + model="policy@*", + ) + + trajectory = group.trajectories[0] + with pytest.raises(ValueError, match="exactly one history"): + trajectory.tokenize(model="policy@*") + tokenized = trajectory.tokenize(model="policy@*", multi_history=True) + assert [history.model for history in tokenized.histories] == [ + "policy@12", + "policy@13", + ] + + +@pytest.mark.parametrize( + ("model", "matches", "misses"), + [ + ("policy@12", ("policy@0", "policy@12"), ("policy@x", "policy@12x")), + ( + "wandb-artifact:///entity/project/run:step12", + ( + "wandb-artifact:///entity/project/run:step0", + "wandb-artifact:///entity/project/run:step12", + ), + ("wandb-artifact:///entity/project/run:stepx",), + ), + ("policy:active", ("policy:active",), ("policy:active2",)), + ("base/model", ("base/model",), ("base/model@1",)), + ], +) +def test_automatic_training_model_selector( + model: str, matches: tuple[str, ...], misses: tuple[str, ...] +) -> None: + selector = automatic_training_model_selector(model) + assert all(selector.matches(candidate) for candidate in matches) + assert not any(selector.matches(candidate) for candidate in misses) + + +def test_automatic_training_model_selector_treats_family_metacharacters_literally() -> ( + None +): + selector = automatic_training_model_selector("policy[blue]*@12") + assert selector.matches("policy[blue]*@13") + assert not selector.matches("policyb@13") + assert not selector.matches("policy[blue]anything@13") + + +def test_automatic_training_model_selector_treats_non_family_metacharacters_literally() -> ( + None +): + selector = automatic_training_model_selector("policy*") + assert selector.matches("policy*") + assert not selector.matches("policy-judge") + + +@pytest.mark.parametrize("automatic", [False, True]) +def test_training_model_selector_rejects_empty_value(automatic: bool) -> None: + with pytest.raises(ValueError, match="cannot be empty"): + if automatic: + automatic_training_model_selector("") + else: + resolve_training_model(_group().trajectories[0], "") + + +def test_automatic_training_selector_rejects_multiple_numeric_steps() -> None: + trajectory = _versioned_group().trajectories[0] + selector = automatic_training_model_selector("policy@12") + with pytest.raises(ValueError, match="exactly one concrete model"): + resolve_training_model(trajectory, selector) + + +def test_public_training_selector_prefers_exact_model_over_glob_interpretation() -> ( + None +): + trajectory = art.Trajectory( + exchanges=art.TrajectoryExchanges( + chat_completions=[ + _exchange("policy*", 2), + _exchange("policyx", 3), + ] + ) + ) + assert resolve_training_model(trajectory, "policy*") == "policy*" + histories = trajectory.histories(model="policy*") + assert [ + history.model + for history in histories + if not isinstance(history, art.LegacyHistory) + ] == ["policy*"] + assert [ + history.model + for history in trajectory.tokenize( + model="policy*", multi_history=True + ).histories + ] == ["policy*"] + + +def test_training_selector_rejects_zero_matches() -> None: + with pytest.raises(ValueError, match="no exchanges"): + resolve_training_model(_group().trajectories[0], "missing*") + + +def test_training_selector_rejects_mixed_protocols() -> None: + start = datetime(2026, 1, 1) + completion = CompletionsExchange( + request=CompletionsRequest(model="policy", prompt="question"), + response=Completion.model_validate( + { + "id": "cmpl", + "object": "text_completion", + "created": 1, + "model": "policy", + "choices": [ + { + "index": 0, + "finish_reason": "stop", + "text": "answer", + } + ], + } + ), + start_time=start, + end_time=start + timedelta(seconds=1), + ) + trajectory = art.Trajectory( + exchanges=art.TrajectoryExchanges( + chat_completions=[_exchange("policy", 2)], + completions=[completion], + ) + ) + with pytest.raises(ValueError, match="mixed protocols"): + resolve_training_model(trajectory, "policy") + + +@pytest.mark.parametrize("output_token_ids", [[2], [2, 3]]) +def test_training_rejects_sampled_token_without_causal_predecessor( + output_token_ids: list[int], + monkeypatch: pytest.MonkeyPatch, +) -> None: + exchange = _empty_prompt_completion_exchange(output_token_ids) + group = art.TrajectoryGroup( + [ + art.Trajectory( + exchanges=art.TrajectoryExchanges(completions=[exchange]), + reward=reward, + ) + for reward in (1.0, 0.0) + ] + ) + tokenized = group.trajectories[0].tokenize() + assert tokenized.token_ids == output_token_ids + assert all(flag & art.TokenFlag.SAMPLED for flag in tokenized.flags) + + weight_writes = 0 + original_setattr = TokenizedResult.__setattr__ + + def tracked_setattr(result: TokenizedResult, name: str, value: object) -> None: + nonlocal weight_writes + if name == "weight": + weight_writes += 1 + original_setattr(result, name, value) + + monkeypatch.setattr(TokenizedResult, "__setattr__", tracked_setattr) + with pytest.raises(ValueError, match="cannot start with a sampled token"): + list( + tokenize_trajectory_groups( + cast(PreTrainedTokenizerBase, _Tokenizer()), + [group], + allow_training_without_logprobs=False, + scale_rewards=False, + model="policy", + ) + ) + assert weight_writes == 0 + + with pytest.raises(ValueError, match="cannot start with a sampled token"): + trajectory_groups_to_datums( + [group], + renderer=None, + tokenizer=None, + normalize_advantages=False, + model="policy", + ) + + +def test_preprocessing_preserves_adjacent_choice_boundaries_for_moe() -> None: + exchanges = [ + _routed_exchange( + prompt_token_ids=[1], + output_token=2, + messages=[{"role": "user", "content": "question"}], + content="first", + ), + _routed_exchange( + prompt_token_ids=[1, 2], + output_token=3, + messages=[ + {"role": "user", "content": "question"}, + {"role": "assistant", "content": "first"}, + {"role": "user", "content": "again"}, + ], + content="second", + ), + ] + group = art.TrajectoryGroup( + trajectories=[ + art.Trajectory( + exchanges=art.TrajectoryExchanges( + chat_completions=exchanges, + ), + reward=reward, + ) + for reward in (1.0, 0.0) + ] + ) + + results = list( + tokenize_trajectory_groups( + cast(PreTrainedTokenizerBase, _Tokenizer()), + [group], + allow_training_without_logprobs=False, + scale_rewards=False, + model="policy", + ) + ) + + assert len(results) == 2 + assert all(result.choice_offsets == [1, 2] for result in results) + assert all(result.moe_routed_experts is not None for result in results) + + +def test_preprocessing_preserves_moe_routes_for_reasoning_stripped_suffix() -> None: + first = _routed_exchange( + prompt_token_ids=[1], + output_token=2, + messages=[{"role": "user", "content": "one"}], + content="first", + ) + first_data = first.response.model_dump(mode="python") + first_data["choices"][0]["message"] = { + "role": "assistant", + "content": "first", + "reasoning": "thought-one", + } + first_data["choices"][0]["token_ids"] = [2, 101, 102, 9] + first_data["choices"][0]["logprobs"]["content"] = [ + { + "token": f"token_id:{token}", + "logprob": -token / 10, + "bytes": [], + "top_logprobs": [], + } + for token in [2, 101, 102, 9] + ] + first.response = ChatCompletion.model_validate(first_data) + first_extra = first.response.choices[0].model_extra + assert first_extra is not None + first_extra[ART_MOE_ROUTING_METADATA_KEY] = { + "prompt_token_ids": [1], + "completion_token_ids": [2, 101, 102, 9], + "routed_experts": np.asarray( + [[[10]], [[20]], [[1010]], [[1020]], [[90]]], dtype=np.int32 + ), + } + + second = _routed_exchange( + prompt_token_ids=[1, 101, 102, 9, 4], + output_token=5, + messages=[ + {"role": "user", "content": "one"}, + {"role": "assistant", "content": "first"}, + {"role": "user", "content": "two"}, + ], + content="second", + ) + second_data = second.response.model_dump(mode="python") + second_data["choices"][0]["message"] = { + "role": "assistant", + "content": "second", + "reasoning": "thought-two", + } + second_data["choices"][0]["token_ids"] = [5, 6] + second_data["choices"][0]["logprobs"]["content"] = [ + { + "token": f"token_id:{token}", + "logprob": -token / 10, + "bytes": [], + "top_logprobs": [], + } + for token in [5, 6] + ] + second.response = ChatCompletion.model_validate(second_data) + second_extra = second.response.choices[0].model_extra + assert second_extra is not None + second_extra[ART_MOE_ROUTING_METADATA_KEY] = { + "prompt_token_ids": [1, 101, 102, 9, 4], + "completion_token_ids": [5, 6], + "routed_experts": np.asarray( + [[[10]], [[1010]], [[1020]], [[90]], [[40]], [[50]], [[60]]], + dtype=np.int32, + ), + } + + class Tokenizer(_Tokenizer): + def __call__(self, text: str, **kwargs: object) -> list[int]: + del kwargs + return { + "one": [1], + "first": [500], + "two": [4], + "thought-two": [5], + "second": [6], + }[text] + + def apply_chat_template( + self, messages: list[dict[str, object]], **kwargs: object + ) -> list[int]: + del kwargs + content = [message.get("content") for message in messages] + return ( + [1, 2, 101, 102, 9] + if content == ["one", "first"] + else [1, 500, 9, 4, 5, 6] + ) + + group = art.TrajectoryGroup( + trajectories=[ + art.Trajectory( + exchanges=art.TrajectoryExchanges( + chat_completions=[first, second], + ), + reward=reward, + ) + for reward in (1.0, 0.0) + ] + ) + + results = list( + tokenize_trajectory_groups( + cast(PreTrainedTokenizerBase, Tokenizer()), + [group], + allow_training_without_logprobs=False, + scale_rewards=False, + shuffle_group_trajectories=False, + drop_zero_advantage_trajectories=False, + model="policy", + ) + ) + + initial = [result for result in results if result.token_ids[1] == 2] + stripped = [result for result in results if result.token_ids[1] == 101] + assert len(initial) == 2 + assert len(stripped) == 2 + assert all(result.choice_offsets == [1] for result in initial) + assert all(result.choice_offsets == [5] for result in stripped) + assert all(result.assistant_mask == [0, 1, 1, 1, 1] for result in initial) + assert all(result.assistant_mask == [0, 0, 0, 0, 0, 1, 1] for result in stripped) + assert all(result.weight == pytest.approx(1 / 6) for result in results) + expected_routes = np.asarray( + [[[10]], [[1010]], [[1020]], [[90]], [[40]], [[50]], [[60]]], + dtype=np.int32, + ) + for result in stripped: + assert isinstance(result.moe_routed_experts, MoeRouteSegments) + assert np.array_equal( + np.concatenate(result.moe_routed_experts.segments), + expected_routes, + ) + + datums = trajectory_groups_to_datums( + [group], + renderer=None, + tokenizer=Tokenizer(), + normalize_advantages=False, + base_model="base/model", + model="policy", + ) + masks = [datum.loss_fn_inputs["mask"].to_torch().tolist() for datum in datums] + assert masks.count([1, 1, 1, 1]) == 2 + assert masks.count([0, 0, 0, 0, 1, 1]) == 2 + + +def test_ambiguous_non_moe_suffix_falls_back_to_sampled_spans() -> None: + first = _exchange("policy", 2) + first_extra = first.response.choices[0].model_extra + assert first_extra is not None + first_extra["token_ids"] = [2, 11] + second = _exchange("policy", 11) + history = ChatCompletionsHistory( + model="policy", + messages=[], + message_sources=[ + ChatCompletionsMessageSource(exchange=first, choice_index=0), + ChatCompletionsMessageSource(exchange=second, choice_index=0), + ], + ) + + assert ( + _chat_choice_trace( + history, + [1, 11, 2, 11], + [ + art.TokenFlag.EXACT, + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + art.TokenFlag.EXACT, + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + ], + ) + is None + ) + + +def test_chat_choice_trace_anchors_retained_suffix_at_its_prompt_boundary() -> None: + first = _exchange("policy", 8) + first_extra = first.response.choices[0].model_extra + assert first_extra is not None + first_extra["prompt_token_ids"] = [1] + first_extra["token_ids"] = [7, 8] + second = _exchange("policy", 8) + second_extra = second.response.choices[0].model_extra + assert second_extra is not None + second_extra["prompt_token_ids"] = [1, 8, 9] + second_extra["token_ids"] = [7, 8] + history = ChatCompletionsHistory( + model="policy", + messages=[], + message_sources=[ + ChatCompletionsMessageSource(exchange=first, choice_index=0), + ChatCompletionsMessageSource(exchange=second, choice_index=0), + ], + ) + + trace = _chat_choice_trace( + history, + [1, 8, 9, 7, 8], + [ + art.TokenFlag.EXACT, + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + art.TokenFlag.EXACT, + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + ], + ) + + assert trace is not None + assert trace.offsets == [1, 3] + assert trace.lengths == [1, 2] + + +def test_chat_choice_trace_does_bounded_work_for_a_retained_suffix() -> None: + class CountingList(list[int]): + slice_reads = 0 + + @overload + def __getitem__(self, key: SupportsIndex, /) -> int: ... + + @overload + def __getitem__(self, key: slice[SupportsIndex | None], /) -> list[int]: ... + + def __getitem__( + self, key: SupportsIndex | slice[SupportsIndex | None], / + ) -> int | list[int]: + if isinstance(key, slice): + self.slice_reads += 1 + return super().__getitem__(key) + + exchange = _exchange("policy", 7) + extra = exchange.response.choices[0].model_extra + assert extra is not None + extra["prompt_token_ids"] = [1] + extra["token_ids"] = [*([8] * 510), 7] + history = ChatCompletionsHistory( + model="policy", + messages=[], + message_sources=[ + ChatCompletionsMessageSource(exchange=exchange, choice_index=0) + ], + ) + token_ids = CountingList([1, 7, *([0] * 510)]) + flags = [ + art.TokenFlag.EXACT, + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + *([art.TokenFlag.EXACT] * 510), + ] + + trace = _chat_choice_trace(history, token_ids, flags) + + assert trace is not None + assert trace.offsets == [1] + assert trace.lengths == [1] + assert token_ids.slice_reads < 10 + + +def test_preprocessing_rejects_partial_choice_evidence_before_moe_routes() -> None: + first = _routed_exchange( + prompt_token_ids=[1], + output_token=2, + messages=[{"role": "user", "content": "question"}], + content="first", + ) + first_extra = first.response.choices[0].model_extra + assert first_extra is not None + first_extra.pop("token_ids") + first_extra.pop(ART_MOE_ROUTING_METADATA_KEY) + second = _routed_exchange( + prompt_token_ids=[1, 2], + output_token=3, + messages=[ + {"role": "user", "content": "question"}, + {"role": "assistant", "content": "first"}, + {"role": "user", "content": "again"}, + ], + content="second", + ) + group = art.TrajectoryGroup( + trajectories=[ + art.Trajectory( + exchanges=art.TrajectoryExchanges(chat_completions=[first, second]), + reward=reward, + ) + for reward in (1.0, 0.0) + ] + ) + + with pytest.raises(RuntimeError, match="every sourced choice"): + list( + tokenize_trajectory_groups( + cast(PreTrainedTokenizerBase, _Tokenizer()), + [group], + allow_training_without_logprobs=False, + scale_rewards=False, + model="policy", + ) + ) diff --git a/tests/unit/test_pipeline_trainer_local_backend.py b/tests/unit/test_pipeline_trainer_local_backend.py index 741b3bed4..e347f22e5 100644 --- a/tests/unit/test_pipeline_trainer_local_backend.py +++ b/tests/unit/test_pipeline_trainer_local_backend.py @@ -753,7 +753,7 @@ def test_local_backend_get_packed_tensors_warns_and_drops_overlong_results( patch( "art.local.backend.tokenize_trajectory_groups", return_value=iter([short_result, long_result]), - ), + ) as tokenize, pytest.warns(UserWarning, match="Dropping 1 tokenized results"), ): packed_tensors = backend._get_packed_tensors( @@ -769,6 +769,10 @@ def test_local_backend_get_packed_tensors_warns_and_drops_overlong_results( assert packed_tensors is not None assert packed_tensors["tokens"].shape == (1, 4) + selector = tokenize.call_args.kwargs["model"] + assert selector.value == f"{model.name}@0" + assert selector.automatic_family == (model.name, "@") + assert selector.allow_glob is False @pytest.mark.asyncio diff --git a/tests/unit/test_tau_bench_client.py b/tests/unit/test_tau_bench_client.py index fd8a06783..ae7fb87e0 100644 --- a/tests/unit/test_tau_bench_client.py +++ b/tests/unit/test_tau_bench_client.py @@ -6,7 +6,8 @@ from typing import Any import httpx -from openai.types.completion_usage import CompletionUsage +from openai import AsyncOpenAI +from openai.types.chat import ChatCompletion import pytest import art @@ -251,16 +252,25 @@ async def delete_environment(self, env_id: str) -> DeleteEnvironmentResponse: class FakeCompletions: async def create(self, **kwargs: Any) -> Any: self.kwargs = kwargs - choice = SimpleNamespace( - message=SimpleNamespace(content="hello", tool_calls=None) - ) - return SimpleNamespace( - choices=[choice], - usage=CompletionUsage( - prompt_tokens=10, - completion_tokens=5, - total_tokens=15, - ), + return ChatCompletion.model_validate( + { + "id": "chat-1", + "object": "chat.completion", + "created": 0, + "model": "default", + "choices": [ + { + "index": 0, + "finish_reason": "stop", + "message": {"role": "assistant", "content": "hello"}, + } + ], + "usage": { + "prompt_tokens": 10, + "completion_tokens": 5, + "total_tokens": 15, + }, + } ) @@ -320,6 +330,167 @@ async def test_rollout_supports_art_model_like_args() -> None: assert trajectory.metrics["num_turns"] == 1 +@pytest.mark.asyncio +async def test_rollout_captures_two_turn_tool_exchange_with_exact_tokens() -> None: + rollout_module = importlib.import_module("art.tau_bench.rollout") + rollout_module.openai_clients.clear() + request_bodies: list[dict[str, Any]] = [] + + class ToolTauBenchClient(FakeTauBenchClient): + async def step_environment( + self, env_id: str, action: str + ) -> StepEnvironmentResponse: + if action == "lookup(key='x')": + return StepEnvironmentResponse( + id=env_id, + observation="tool: result", + reward=0.25, + terminated=False, + truncated=False, + info={}, + ) + assert action == "hello" + return StepEnvironmentResponse( + id=env_id, + observation="user: done", + reward=0.75, + terminated=True, + truncated=False, + info={"user_message_cost": 0.25}, + ) + + async def handler(request: httpx.Request) -> httpx.Response: + request_bodies.append(json.loads(request.content)) + if len(request_bodies) == 1: + return httpx.Response( + 200, + json={ + "id": "chat-tool", + "object": "chat.completion", + "created": 0, + "model": "default", + "prompt_token_ids": [10, 11], + "choices": [ + { + "index": 0, + "finish_reason": "tool_calls", + "message": { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call-1", + "type": "function", + "function": { + "name": "lookup", + "arguments": '{"key":"x"}', + }, + } + ], + }, + "token_ids": [12], + "logprobs": { + "content": [ + { + "token": "token_id:12", + "logprob": -0.25, + "bytes": None, + "top_logprobs": [], + } + ] + }, + } + ], + "usage": { + "prompt_tokens": 2, + "completion_tokens": 1, + "total_tokens": 3, + }, + }, + ) + return httpx.Response( + 200, + json={ + "id": "chat-exact", + "object": "chat.completion", + "created": 0, + "model": "default", + "prompt_token_ids": [10, 11, 12, 13], + "choices": [ + { + "index": 0, + "finish_reason": "stop", + "message": {"role": "assistant", "content": "hello"}, + "token_ids": [14], + "logprobs": { + "content": [ + { + "token": "token_id:14", + "logprob": -0.5, + "bytes": [104, 101, 108, 108, 111], + "top_logprobs": [], + } + ] + }, + } + ], + "usage": { + "prompt_tokens": 4, + "completion_tokens": 1, + "total_tokens": 5, + }, + }, + ) + + http_client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + openai_client = AsyncOpenAI( + api_key="model-key", + base_url="http://model.test/v1", + http_client=http_client, + ) + rollout_module.openai_clients[("http://model.test/v1", "model-key")] = openai_client + try: + trajectory = await rollout_module.rollout( + Scenario(domain="banking_knowledge", task=Task(id="task_001")), + "http://model.test/v1", + "model-key", + "default", + client=ToolTauBenchClient(), + max_turns=2, + ) + finally: + await openai_client.close() + await http_client.aclose() + rollout_module.openai_clients.clear() + + assert request_bodies[0]["messages"] == [ + {"role": "system", "content": "policy"}, + {"role": "user", "content": "hello"}, + ] + assert request_bodies[1]["messages"][2]["tool_calls"][0]["id"] == "call-1" + assert request_bodies[1]["messages"][3] == { + "role": "tool", + "content": "result", + "tool_call_id": "call-1", + } + assert len(trajectory.exchanges.chat_completions) == 2 + assert not trajectory.messages_and_choices + assert trajectory.tools is None + restored = art.Trajectory.model_validate_json(trajectory.model_dump_json()) + tokenized = restored.tokenize() + assert tokenized.token_ids == [10, 11, 12, 13, 14] + assert tokenized.logprobs[2] == -0.25 + assert tokenized.logprobs[3] != tokenized.logprobs[3] + assert tokenized.logprobs[4] == -0.5 + assert tokenized.flags == [ + art.TokenFlag.EXACT, + art.TokenFlag.EXACT, + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + art.TokenFlag.EXACT, + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + ] + + class FakeBadRequestError(Exception): def __init__(self, message: str) -> None: self.message = message @@ -382,16 +553,25 @@ async def test_rollout_stops_on_max_tokens_bad_request( class NearContextLimitCompletions: async def create(self, **kwargs: Any) -> Any: - choice = SimpleNamespace( - message=SimpleNamespace(content="hello", tool_calls=None) - ) - return SimpleNamespace( - choices=[choice], - usage=CompletionUsage( - prompt_tokens=32_000, - completion_tokens=700, - total_tokens=32_700, - ), + return ChatCompletion.model_validate( + { + "id": "chat-near-limit", + "object": "chat.completion", + "created": 0, + "model": "default", + "choices": [ + { + "index": 0, + "finish_reason": "stop", + "message": {"role": "assistant", "content": "hello"}, + } + ], + "usage": { + "prompt_tokens": 32_000, + "completion_tokens": 700, + "total_tokens": 32_700, + }, + } ) diff --git a/tests/unit/test_tinker_native_exchanges.py b/tests/unit/test_tinker_native_exchanges.py index d1bdcb8a4..3e7dd4359 100644 --- a/tests/unit/test_tinker_native_exchanges.py +++ b/tests/unit/test_tinker_native_exchanges.py @@ -1,7 +1,7 @@ from datetime import datetime from typing import Any -from openai.types.chat import ChatCompletion +from openai.types.chat import ChatCompletion, ChatCompletionMessageParam from openai.types.chat.chat_completion import Choice import pytest @@ -16,6 +16,16 @@ def _exchange( prompt: list[int], output: list[int], *, logprobs: bool = True ) -> ChatCompletionsExchange: + messages: list[ChatCompletionMessageParam] = [ + {"role": "user", "content": "question"} + ] + if len(prompt) > 1: + messages.extend( + [ + {"role": "assistant", "content": "answer"}, + {"role": "user", "content": "next"}, + ] + ) response = ChatCompletion.model_validate( { "id": "chat-1", @@ -48,7 +58,7 @@ def _exchange( } ) return ChatCompletionsExchange( - request=ChatCompletionsRequest(model="test/model", messages=[]), + request=ChatCompletionsRequest(model="test/model", messages=messages), response=response, start_time=datetime.now(), end_time=datetime.now(), diff --git a/tests/unit/test_trajectory_parquet.py b/tests/unit/test_trajectory_parquet.py index e766239b2..a0898a14f 100644 --- a/tests/unit/test_trajectory_parquet.py +++ b/tests/unit/test_trajectory_parquet.py @@ -35,7 +35,7 @@ ) import pytest -from art import History, Trajectory, TrajectoryGroup +from art import LegacyHistory, Trajectory, TrajectoryGroup from art.types import MessageOrChoice from art.utils.trajectory_logging import ( read_trajectory_groups_parquet, @@ -154,7 +154,19 @@ def _exchange_trajectory() -> Trajectory: "parallel_tool_calls": True, "tool_choice": "auto", "tools": [], - "raw_output_tokens": [{"token_id": 20, "logprob": -0.3}], + "token_generations": [ + { + "prompt_token_ids": [10], + "output_tokens": [ + { + "token_id": 20, + "logprob": -0.3, + "text": "done", + } + ], + "output_indices": [0], + } + ], }, "start_time": "2026-01-01T00:00:04", "end_time": "2026-01-01T00:00:05", @@ -295,7 +307,7 @@ def test_legacy_serializer_round_trips_complete_legacy_trajectory() -> None: [ Trajectory( messages_and_choices=[choice], - additional_histories=[History(messages_and_choices=[choice])], + additional_histories=[LegacyHistory(messages_and_choices=[choice])], reward=1.0, initial_policy_version=3, final_policy_version=4, @@ -513,7 +525,9 @@ def test_complete_models_round_trip(self, tmp_path: Path) -> None: } ], additional_histories=[ - History(messages_and_choices=[{"role": "user", "content": "alternate"}]) + LegacyHistory( + messages_and_choices=[{"role": "user", "content": "alternate"}] + ) ], initial_policy_version=7, final_policy_version=8, diff --git a/tests/unit/trajectories/test_capture.py b/tests/unit/trajectories/test_capture.py index c28c96157..7a2d2a647 100644 --- a/tests/unit/trajectories/test_capture.py +++ b/tests/unit/trajectories/test_capture.py @@ -4,16 +4,18 @@ from collections.abc import AsyncGenerator, AsyncIterator, Generator import copy from datetime import datetime, timedelta +import gzip import json from typing import Any, cast from unittest.mock import Mock +import zlib import aiohttp from aiohttp import web from anthropic import AsyncAnthropic from anthropic.types import TextBlock import httpx -from openai import AsyncOpenAI +from openai import AsyncOpenAI, OpenAI import pytest import pytest_asyncio import requests @@ -111,7 +113,13 @@ "output_tokens_details": {"reasoning_tokens": 0}, "total_tokens": 5, }, - "raw_output_tokens": [{"token_id": 2, "logprob": -0.2}], + "token_generations": [ + { + "prompt_token_ids": [1], + "output_tokens": [{"token_id": 2, "logprob": -0.2, "text": "hello"}], + "output_indices": [0], + } + ], } MESSAGE: dict[str, Any] = { @@ -123,11 +131,39 @@ "stop_reason": "end_turn", "stop_sequence": None, "usage": {"input_tokens": 1, "output_tokens": 1}, + "prompt_token_ids": [1], "token_ids": [2], "logprobs": [-0.2], } +class _SyncChunks(httpx.SyncByteStream): + def __init__(self, body: bytes) -> None: + self.body = body + + def __iter__(self) -> Generator[bytes, None, None]: + for index in range(0, len(self.body), 3): + yield self.body[index : index + 3] + + +class _AsyncChunks(httpx.AsyncByteStream): + def __init__(self, body: bytes) -> None: + self.body = body + + async def __aiter__(self) -> AsyncIterator[bytes]: + for index in range(0, len(self.body), 3): + yield self.body[index : index + 3] + + +def _encoded(body: bytes, encoding: str) -> bytes: + if encoding == "gzip": + return gzip.compress(body) + if encoding == "deflate": + return zlib.compress(body) + brotli = pytest.importorskip("brotli") + return cast(bytes, brotli.compress(body)) + + @pytest_asyncio.fixture async def endpoint_server(unused_tcp_port: int) -> AsyncIterator[str]: async def handler(request: web.Request) -> web.StreamResponse: @@ -172,7 +208,7 @@ async def handler(request: web.Request) -> web.StreamResponse: async def test_contexts_are_nested_and_task_local() -> None: assert art.current_trajectory() is None with art.Trajectory() as outer: - assert art.current_trajectory(required=True) is outer + assert art.current_trajectory(require=True) is outer with art.Trajectory() as inner: assert art.current_trajectory() is inner assert art.current_trajectory() is outer @@ -187,7 +223,7 @@ async def child() -> art.Trajectory: assert first is not second assert art.current_trajectory() is None with pytest.raises(RuntimeError, match="No trajectory"): - art.current_trajectory(required=True) + art.current_trajectory(require=True) async def test_group_context_and_async_helpers() -> None: @@ -303,6 +339,233 @@ def requests_stream() -> None: ) +@pytest.mark.parametrize("encoding", ["gzip", "deflate", "br"]) +@pytest.mark.parametrize("mode", ["raw", "bytes", "lines"]) +def test_httpx_sync_stream_consumption_captures_decoded_body_once( + encoding: str, mode: str +) -> None: + body = _streaming_chat_body() + compressed = _encoded(body, encoding) + + def response(_: httpx.Request) -> httpx.Response: + return httpx.Response( + 200, + headers={"content-encoding": encoding}, + stream=_SyncChunks(compressed), + ) + + with art.Trajectory() as trajectory: + with httpx.Client(transport=httpx.MockTransport(response)) as client: + with client.stream( + "POST", + "https://example.test/v1/chat/completions", + json={"model": "test/model", "messages": [], "stream": True}, + ) as result: + if mode == "raw": + assert b"".join(result.iter_raw()) == compressed + elif mode == "bytes": + assert b"".join(result.iter_bytes()) == body + else: + list(result.iter_lines()) + + assert len(trajectory.exchanges.chat_completions) == 1 + + +@pytest.mark.parametrize("encoding", ["gzip", "deflate", "br"]) +@pytest.mark.parametrize("mode", ["raw", "bytes", "lines"]) +async def test_httpx_async_stream_consumption_captures_decoded_body_once( + encoding: str, mode: str +) -> None: + body = _streaming_chat_body() + compressed = _encoded(body, encoding) + + async def response(_: httpx.Request) -> httpx.Response: + return httpx.Response( + 200, + headers={"content-encoding": encoding}, + stream=_AsyncChunks(compressed), + ) + + with art.Trajectory() as trajectory: + async with httpx.AsyncClient(transport=httpx.MockTransport(response)) as client: + async with client.stream( + "POST", + "https://example.test/v1/chat/completions", + json={"model": "test/model", "messages": [], "stream": True}, + ) as result: + if mode == "raw": + assert ( + b"".join([chunk async for chunk in result.aiter_raw()]) + == compressed + ) + elif mode == "bytes": + assert ( + b"".join([chunk async for chunk in result.aiter_bytes()]) + == body + ) + else: + _ = [line async for line in result.aiter_lines()] + + assert len(trajectory.exchanges.chat_completions) == 1 + + +@pytest.mark.parametrize("encoding", ["gzip", "deflate", "br"]) +async def test_httpx_terminal_event_captures_before_response_close( + encoding: str, +) -> None: + body = _streaming_chat_body() + compressed = _encoded(body, encoding) + + async def response(_: httpx.Request) -> httpx.Response: + return httpx.Response( + 200, + headers={ + "content-encoding": encoding, + "content-type": "text/event-stream", + }, + stream=_AsyncChunks(compressed), + ) + + client = httpx.AsyncClient(transport=httpx.MockTransport(response)) + request = client.build_request( + "POST", + "https://example.test/v1/chat/completions", + json={"model": "test/model", "messages": [], "stream": True}, + ) + with art.Trajectory() as trajectory: + result = await client.send(request, stream=True) + iterator = cast(AsyncGenerator[bytes, None], result.aiter_bytes()) + received = bytearray() + async for chunk in iterator: + received.extend(chunk) + if b"data: [DONE]\n\n" in received: + break + assert len(trajectory.exchanges.chat_completions) == 1 + + assert not result.is_closed + await iterator.aclose() + await result.aclose() + await client.aclose() + + +def test_httpx_raw_capture_failure_does_not_change_user_stream() -> None: + malformed = b"not a gzip stream" + + def response(_: httpx.Request) -> httpx.Response: + return httpx.Response( + 200, + headers={"content-encoding": "gzip"}, + stream=_SyncChunks(malformed), + ) + + with art.Trajectory() as trajectory: + with httpx.Client(transport=httpx.MockTransport(response)) as client: + with client.stream( + "POST", + "https://example.test/v1/chat/completions", + json={"model": "test/model", "messages": [], "stream": True}, + ) as result: + assert b"".join(result.iter_raw()) == malformed + + assert not trajectory.exchanges + + +async def test_httpx_async_raw_capture_failure_does_not_change_user_stream() -> None: + malformed = b"not a gzip stream" + + async def response(_: httpx.Request) -> httpx.Response: + return httpx.Response( + 200, + headers={"content-encoding": "gzip"}, + stream=_AsyncChunks(malformed), + ) + + with art.Trajectory() as trajectory: + async with httpx.AsyncClient(transport=httpx.MockTransport(response)) as client: + async with client.stream( + "POST", + "https://example.test/v1/chat/completions", + json={"model": "test/model", "messages": [], "stream": True}, + ) as result: + assert ( + b"".join([chunk async for chunk in result.aiter_raw()]) == malformed + ) + + assert not trajectory.exchanges + + +def _streaming_chat_body() -> bytes: + return _sse( + [ + ( + None, + { + "id": "chatcmpl-1", + "object": "chat.completion.chunk", + "created": 1, + "model": "test/model", + "choices": [ + { + "index": 0, + "delta": {"role": "assistant", "content": "hello"}, + "finish_reason": "stop", + } + ], + }, + ), + (None, "[DONE]"), + ] + ) + + +def test_httpx_abandoned_compressed_raw_stream_is_excluded() -> None: + compressed = gzip.compress(_streaming_chat_body()) + + def response(_: httpx.Request) -> httpx.Response: + return httpx.Response( + 200, + headers={"content-encoding": "gzip"}, + stream=_SyncChunks(compressed), + ) + + with art.Trajectory() as trajectory: + with httpx.Client(transport=httpx.MockTransport(response)) as client: + with client.stream( + "POST", + "https://example.test/v1/chat/completions", + json={"model": "test/model", "messages": [], "stream": True}, + ) as result: + iterator = cast(Generator[bytes, None, None], result.iter_raw()) + next(iterator) + iterator.close() + + assert not trajectory.exchanges + + +async def test_httpx_abandoned_compressed_async_raw_stream_is_excluded() -> None: + compressed = gzip.compress(_streaming_chat_body()) + + async def response(_: httpx.Request) -> httpx.Response: + return httpx.Response( + 200, + headers={"content-encoding": "gzip"}, + stream=_AsyncChunks(compressed), + ) + + with art.Trajectory() as trajectory: + async with httpx.AsyncClient(transport=httpx.MockTransport(response)) as client: + async with client.stream( + "POST", + "https://example.test/v1/chat/completions", + json={"model": "test/model", "messages": [], "stream": True}, + ) as result: + iterator = cast(AsyncGenerator[bytes, None], result.aiter_raw()) + await anext(iterator) + await iterator.aclose() + + assert not trajectory.exchanges + + async def test_aiohttp_capture_covers_stream_reader_consumption_methods( endpoint_server: str, ) -> None: @@ -349,6 +612,28 @@ def consume(trajectory: art.Trajectory) -> None: assert len(trajectory.exchanges.chat_completions) == 1 +async def test_requests_decode_unicode_preserves_string_chunks( + endpoint_server: str, +) -> None: + body = {"model": "test/model", "messages": [{"role": "user", "content": "hi"}]} + + def consume() -> list[str | bytes]: + with requests.post( + f"{endpoint_server}/chat/completions", + json=body, + stream=True, + timeout=5, + ) as response: + return list(response.iter_content(chunk_size=5, decode_unicode=True)) + + with art.Trajectory() as trajectory: + chunks = await asyncio.to_thread(consume) + + assert chunks + assert all(isinstance(chunk, str) for chunk in chunks) + assert len(trajectory.exchanges.chat_completions) == 1 + + async def test_native_openai_and_anthropic_sdks(endpoint_server: str) -> None: openai = AsyncOpenAI(base_url=endpoint_server, api_key="test") anthropic = AsyncAnthropic( @@ -374,6 +659,136 @@ async def test_native_openai_and_anthropic_sdks(endpoint_server: str) -> None: assert len(trajectory.exchanges.messages) == 1 +async def test_native_openai_chat_stream_captures_at_done_event() -> None: + async def response(_: httpx.Request) -> httpx.Response: + return httpx.Response( + 200, + headers={"content-type": "text/event-stream"}, + stream=_AsyncChunks(_streaming_chat_body()), + ) + + http_client = httpx.AsyncClient(transport=httpx.MockTransport(response)) + client = AsyncOpenAI( + base_url="https://example.test/v1", + api_key="test", + http_client=http_client, + ) + with art.Trajectory() as trajectory: + stream = await client.chat.completions.create( + model="test/model", + messages=[], + stream=True, + ) + chunks = [chunk async for chunk in stream] + assert len(trajectory.exchanges.chat_completions) == 1 + + assert len(chunks) == 1 + assert chunks[0].choices[0].delta.content == "hello" + await stream.close() + assert len(trajectory.exchanges.chat_completions) == 1 + await client.close() + + +def test_native_openai_preloaded_chat_stream_captures_once() -> None: + def response(_: httpx.Request) -> httpx.Response: + return httpx.Response( + 200, + headers={"content-type": "text/event-stream"}, + content=_streaming_chat_body(), + ) + + http_client = httpx.Client(transport=httpx.MockTransport(response)) + client = OpenAI( + base_url="https://example.test/v1", + api_key="test", + http_client=http_client, + ) + with art.Trajectory() as trajectory: + stream = client.chat.completions.create( + model="test/model", + messages=[], + stream=True, + ) + assert len(trajectory.exchanges.chat_completions) == 1 + chunks = list(stream) + assert len(trajectory.exchanges.chat_completions) == 1 + + assert len(chunks) == 1 + assert chunks[0].choices[0].delta.content == "hello" + stream.close() + assert len(trajectory.exchanges.chat_completions) == 1 + client.close() + + +async def test_native_async_openai_preloaded_chat_stream_captures_once() -> None: + async def response(_: httpx.Request) -> httpx.Response: + return httpx.Response( + 200, + headers={"content-type": "text/event-stream"}, + content=_streaming_chat_body(), + ) + + http_client = httpx.AsyncClient(transport=httpx.MockTransport(response)) + client = AsyncOpenAI( + base_url="https://example.test/v1", + api_key="test", + http_client=http_client, + ) + with art.Trajectory() as trajectory: + stream = await client.chat.completions.create( + model="test/model", + messages=[], + stream=True, + ) + assert len(trajectory.exchanges.chat_completions) == 1 + chunks = [chunk async for chunk in stream] + assert len(trajectory.exchanges.chat_completions) == 1 + + assert len(chunks) == 1 + assert chunks[0].choices[0].delta.content == "hello" + await stream.close() + assert len(trajectory.exchanges.chat_completions) == 1 + await client.close() + + +def test_preloaded_malformed_stream_is_excluded() -> None: + def response(_: httpx.Request) -> httpx.Response: + return httpx.Response( + 200, + headers={"content-type": "text/event-stream"}, + content=_sse([(None, "corrupt"), (None, "[DONE]")]), + ) + + with art.Trajectory() as trajectory: + with httpx.Client(transport=httpx.MockTransport(response)) as client: + with client.stream( + "POST", + "https://example.test/v1/chat/completions", + json={"model": "test/model", "messages": [], "stream": True}, + ) as result: + assert list(result.iter_bytes()) + + assert not trajectory.exchanges + + +def test_stream_terminal_without_sse_boundary_is_excluded() -> None: + body = _streaming_chat_body()[:-1] + with art.Trajectory() as trajectory: + state, token = begin( + "POST", + "https://example.test/v1/chat/completions", + {"model": "test/model", "messages": [], "stream": True}, + ) + reset(token) + assert state is not None + state.status_code = 200 + state.add(body) + assert not state.captured + state.finish() + + assert not trajectory.exchanges + + async def test_failed_and_incomplete_calls_are_excluded(endpoint_server: str) -> None: async with httpx.AsyncClient() as client: with art.Trajectory() as trajectory: @@ -634,6 +1049,7 @@ def test_all_streaming_protocols_reconstruct_final_responses() -> None: } response_event = {"type": "response.completed", "response": RESPONSE} message_events = [ + ("ping", {"type": "ping"}), ( "message_start", { @@ -643,6 +1059,7 @@ def test_all_streaming_protocols_reconstruct_final_responses() -> None: "content": [], "stop_reason": None, "usage": {"input_tokens": 1, "output_tokens": 0}, + "prompt_token_ids": [1], }, }, ), @@ -660,6 +1077,8 @@ def test_all_streaming_protocols_reconstruct_final_responses() -> None: "type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "hello"}, + "token_ids": [2], + "logprobs": [-0.2], }, ), ("content_block_stop", {"type": "content_block_stop", "index": 0}), @@ -669,6 +1088,7 @@ def test_all_streaming_protocols_reconstruct_final_responses() -> None: "type": "message_delta", "delta": {"stop_reason": "end_turn", "stop_sequence": None}, "usage": {"output_tokens": 1}, + "prompt_token_ids": [1], "token_ids": [2], "logprobs": [-0.2], }, @@ -710,8 +1130,143 @@ def test_all_streaming_protocols_reconstruct_final_responses() -> None: content = exchange.response.content[0] assert isinstance(content, TextBlock) assert content.text == "hello" + assert content.model_extra is not None + assert content.model_extra["token_ids"] == [2] + assert content.model_extra["logprobs"] == [-0.2] assert getattr(exchange.response, "token_ids") == [2] assert getattr(exchange.response, "logprobs") == [-0.2] + assert getattr(exchange.response, "prompt_token_ids") == [1] + + with art.Trajectory() as trajectory: + state, token = begin( + "POST", + f"https://example.test/v1/{endpoint.replace('_', '/')}", + request, + ) + reset(token) + assert state is not None + state.status_code = 200 + for byte in body[:-1]: + state.add(bytes([byte])) + assert not state.captured + state.add(body[-1:]) + assert state.captured + + assert ( + sum( + len(exchanges) + for exchanges in ( + trajectory.exchanges.chat_completions, + trajectory.exchanges.completions, + trajectory.exchanges.responses, + trajectory.exchanges.messages, + ) + ) + == 1 + ) + + +def test_streaming_messages_error_event_is_rejected_and_not_captured() -> None: + body = _sse( + [ + ( + "error", + { + "type": "error", + "error": {"type": "api_error", "message": "failed"}, + }, + ) + ] + ) + request = {"model": "test/model", "messages": [], "stream": True} + with pytest.raises(ValueError, match="returned an error event"): + build_exchange( + "messages", + request, + body, + start_time=datetime.now(), + end_time=datetime.now(), + ) + + with art.Trajectory() as trajectory: + state, token = begin("POST", "https://example.test/v1/messages", request) + reset(token) + assert state is not None + state.status_code = 200 + state.add(body) + state.finish() + + assert not trajectory.exchanges + + +@pytest.mark.parametrize( + ("field", "value"), + [ + ("prompt_token_ids", [True]), + ("prompt_token_ids", [-1]), + ("token_ids", [False]), + ("token_ids", [-1]), + ("logprobs", [True]), + ("logprobs", [float("nan")]), + ("logprobs", [float("inf")]), + ], +) +def test_streaming_messages_reject_malformed_token_metadata( + field: str, value: list[object] +) -> None: + delta = { + "type": "message_delta", + "delta": {"stop_reason": "end_turn", "stop_sequence": None}, + "usage": {"output_tokens": 1}, + "prompt_token_ids": [1], + "token_ids": [2], + "logprobs": [-0.2], + field: value, + } + body = _sse( + [ + ( + "message_start", + { + "type": "message_start", + "message": { + **MESSAGE, + "content": [], + "stop_reason": None, + "usage": {"input_tokens": 1, "output_tokens": 0}, + }, + }, + ), + ( + "content_block_start", + { + "type": "content_block_start", + "index": 0, + "content_block": {"type": "text", "text": "", "citations": None}, + }, + ), + ( + "content_block_delta", + { + "type": "content_block_delta", + "index": 0, + "delta": {"type": "text_delta", "text": "hello"}, + }, + ), + ("content_block_stop", {"type": "content_block_stop", "index": 0}), + ("message_delta", delta), + ("message_stop", {"type": "message_stop"}), + ] + ) + + with pytest.raises(ValueError, match="must contain"): + build_exchange( + "messages", + {"model": "test/model", "messages": [], "stream": True}, + body, + start_time=datetime.now(), + end_time=datetime.now(), + ) def test_streaming_chat_choices_are_accumulated_by_index() -> None: @@ -757,6 +1312,45 @@ def chunk(index: int, content: str) -> dict[str, Any]: ] +def test_streaming_chat_preserves_reasoning_fields() -> None: + now = datetime.now() + + def chunk(**delta: str) -> dict[str, Any]: + return { + "id": "chatcmpl-1", + "object": "chat.completion.chunk", + "created": 1, + "model": "test/model", + "choices": [ + { + "index": 0, + "delta": {"role": "assistant", **delta}, + "finish_reason": None, + "logprobs": None, + } + ], + } + + exchange = build_exchange( + "chat_completions", + {"model": "test/model", "messages": [], "stream": True}, + _sse( + [ + (None, chunk(reasoning="r1", reasoning_content="c1")), + (None, chunk(reasoning="r2", reasoning_content="c2")), + (None, "[DONE]"), + ] + ), + start_time=now, + end_time=now + timedelta(seconds=1), + ) + + assert isinstance(exchange, ChatCompletionsExchange) + message = exchange.response.choices[0].message + assert getattr(message, "reasoning") == "r1r2" + assert getattr(message, "reasoning_content") == "c1c2" + + def test_streaming_chat_ignores_keepalives_and_azure_prologue() -> None: chunk = { "id": "chatcmpl-1", diff --git a/tests/unit/trajectories/test_history.py b/tests/unit/trajectories/test_history.py index 809fa8196..a79deffad 100644 --- a/tests/unit/trajectories/test_history.py +++ b/tests/unit/trajectories/test_history.py @@ -1,5 +1,8 @@ from datetime import datetime, timedelta import importlib +from statistics import median +from time import perf_counter +from typing import Any, cast from anthropic.types import Message from openai.types import Completion @@ -12,9 +15,11 @@ ChatCompletionsExchange, CompletionsExchange, MessagesExchange, + MessagesRequest, ResponsesExchange, TrajectoryExchanges, ) +from art.types import Message as ChatMessage def _times(offset: int = 0) -> tuple[datetime, datetime]: @@ -52,6 +57,16 @@ def _chat( ) +def _growing_chat_trajectory(turn_count: int) -> art.Trajectory: + exchanges: list[ChatCompletionsExchange] = [] + messages: list[dict[str, object]] = [] + for index in range(turn_count): + messages.append({"role": "user", "content": f"question {index}"}) + exchanges.append(_chat(list(messages), f"answer {index}", offset=index)) + messages.append({"role": "assistant", "content": f"answer {index}"}) + return art.Trajectory(exchanges=TrajectoryExchanges(chat_completions=exchanges)) + + def _completion( prompt: list[int], output: list[int], *, offset: int = 0 ) -> CompletionsExchange: @@ -199,30 +214,347 @@ def test_chat_history_resolves_one_model_and_append_only_sequence() -> None: "user", "assistant", ] + assert [source.exchange for source in history.message_sources if source] == [ + first, + first, + second, + second, + ] + assert history.messages is not first.request["messages"] assert trajectory.chat_completions_history(model="test/model") == history second.request["cache_salt"] = "new-cache" - with pytest.raises(ValueError, match="different cache_salt"): - trajectory.chat_completions_history(model="test/model") + assert trajectory.chat_completions_history(model="test/model") == history second.request.pop("cache_salt") second.request["messages"] = [{"role": "user", "content": "branch"}] - with pytest.raises(ValueError, match="append-only"): + assert len(trajectory.chat_completions_histories(model="test/model")) == 2 + with pytest.raises(ValueError, match="exactly one history"): trajectory.chat_completions_history(model="test/model") +def test_chat_projection_keys_each_captured_message_once( + monkeypatch: pytest.MonkeyPatch, +) -> None: + history_module = importlib.import_module("art.trajectories._history") + original = history_module._chat_message_key + calls = 0 + + def count(message: ChatMessage, *, visible_only: bool = False) -> str: + nonlocal calls + calls += 1 + return original(message, visible_only=visible_only) + + monkeypatch.setattr(history_module, "_chat_message_key", count) + trajectory = _growing_chat_trajectory(16) + + trajectory.chat_completions_history() + + expected = sum( + len(exchange.request["messages"]) + len(exchange.response.choices) + for exchange in trajectory.exchanges.chat_completions + ) + assert calls == expected + + +def test_chat_projection_scales_with_captured_messages() -> None: + measurements: list[tuple[int, float, int]] = [] + for turn_count in (32, 64, 128): + trajectory = _growing_chat_trajectory(turn_count) + captured_bytes = sum( + len(exchange.model_dump_json().encode()) + for exchange in trajectory.exchanges.chat_completions + ) + trajectory.chat_completions_history() + samples: list[float] = [] + for _ in range(5): + started = perf_counter() + trajectory.chat_completions_history() + samples.append(perf_counter() - started) + elapsed = median(samples) + measurements.append((turn_count, elapsed, captured_bytes)) + + # A growing transcript contains O(turns²) serialized evidence: doubling the + # turn count roughly quadruples the bytes that projection must validate. + # Preserve the turn-count curve as diagnostics, but gate near-linear cost in + # the actual input size rather than imposing an impossible per-turn ratio. + normalized = [ + elapsed / captured_bytes for _, elapsed, captured_bytes in measurements + ] + assert normalized[1] < normalized[0] * 2, measurements + assert normalized[2] < normalized[1] * 2, measurements + + +def test_model_patterns_select_matching_histories_only() -> None: + policy_12 = _chat( + [{"role": "user", "content": "one"}], + "first", + model="policy@12", + ) + judge = _chat( + [{"role": "user", "content": "judge"}], + "score", + model="judge@4", + offset=1, + ) + policy_13 = _chat( + [{"role": "user", "content": "two"}], + "second", + model="policy@13", + offset=2, + ) + trajectory = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[policy_12, judge, policy_13]) + ) + + histories = trajectory.chat_completions_histories(model="policy@*") + assert [history.model for history in histories] == ["policy@12", "policy@13"] + generic_histories = trajectory.histories(model="policy@*") + assert all(isinstance(history, art.History) for history in generic_histories) + assert [cast(art.History, history).model for history in generic_histories] == [ + "policy@12", + "policy@13", + ] + assert trajectory.chat_completions_history(model="policy@12").model == "policy@12" + with pytest.raises(ValueError, match="exactly one history"): + trajectory.history(model="policy@*") + with pytest.raises(ValueError, match="no Chat Completions exchanges"): + trajectory.chat_completions_histories(model="foreign@*") + + +def test_chat_history_preserves_provider_specific_nested_fields() -> None: + exchange = _chat( + [ + { + "role": "user", + "content": [ + { + "type": "text", + "text": "one", + "cache_control": {"type": "ephemeral"}, + } + ], + } + ], + "first", + ) + + trajectory = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[exchange]) + ) + history = trajectory.chat_completions_history() + + content = cast(list[dict[str, object]], history.messages[0]["content"]) + assert content[0]["cache_control"] == {"type": "ephemeral"} + dumped = trajectory.model_dump(mode="json", warnings="error") + dumped_content = dumped["exchanges"]["chat_completions"][0]["request"]["messages"][ + 0 + ]["content"] + assert dumped_content[0]["cache_control"] == {"type": "ephemeral"} + + +def test_chat_choices_branch_and_identical_continuation_uses_first_choice() -> None: + first = _chat([{"role": "user", "content": "one"}], "same") + response = first.response.model_dump(mode="python") + second_choice = dict(response["choices"][0]) + second_choice["index"] = 1 + response["choices"].append(second_choice) + first.response = ChatCompletion.model_validate(response) + continuation = _chat( + [ + {"role": "user", "content": "one"}, + {"role": "assistant", "content": "same"}, + {"role": "user", "content": "two"}, + ], + "continued", + offset=1, + ) + trajectory = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[first, continuation]) + ) + + histories = trajectory.chat_completions_histories() + + assert len(histories) == 2 + assert [len(history.messages) for history in histories] == [4, 2] + assert histories[0].message_sources[1] is not None + assert histories[0].message_sources[1].choice_index == 0 + assert histories[1].message_sources[1] is not None + assert histories[1].message_sources[1].choice_index == 1 + + +def test_chat_history_normalizes_empty_response_only_fields() -> None: + first = _chat([{"role": "user", "content": "one"}], "first") + data = first.response.model_dump(mode="python") + data["choices"][0]["message"]["tool_calls"] = [] + data["choices"][0]["message"]["annotations"] = [] + first.response = ChatCompletion.model_validate(data) + second = _chat( + [ + {"role": "user", "content": "one"}, + {"role": "assistant", "content": "first"}, + {"role": "user", "content": "two"}, + ], + "second", + offset=1, + ) + + history = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[first, second]) + ).chat_completions_history() + + assert len(history.messages) == 4 + assert "tool_calls" not in history.messages[1] + assert "annotations" not in history.messages[1] + assert history.message_sources[1] is not None + assert history.message_sources[1].exchange is first + + +def test_chat_history_normalizes_assistant_missing_content_to_empty() -> None: + first = _chat([{"role": "user", "content": ""}], "") + data = first.response.model_dump(mode="python") + data["choices"][0]["message"] = { + "role": "assistant", + "content": None, + "annotations": [], + } + first.response = ChatCompletion.model_validate(data) + second = _chat( + [ + {"role": "user", "content": ""}, + {"role": "assistant", "content": ""}, + {"role": "user", "content": "next"}, + ], + "second", + offset=1, + ) + + history = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[first, second]) + ).chat_completions_history() + + assert [message["content"] for message in history.messages] == [ + "", + "", + "next", + "second", + ] + assert "annotations" not in history.messages[1] + assert history.message_sources[1] is not None + assert history.message_sources[1].exchange is first + + +@pytest.mark.parametrize("indices", [[], [0, 0]]) +def test_chat_history_rejects_missing_or_duplicate_choice_indices( + indices: list[int], +) -> None: + exchange = _chat([{"role": "user", "content": "one"}], "first") + data = exchange.response.model_dump(mode="python") + choice = data["choices"][0] + data["choices"] = [{**choice, "index": index} for index in indices] + exchange.response = ChatCompletion.model_validate(data) + + with pytest.raises(ValueError, match="choices|choice indices"): + art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[exchange]) + ).chat_completions_histories() + + +def test_same_content_seeded_inputs_remain_request_sourced() -> None: + first_chat = _chat([{"role": "user", "content": "prompt"}], "same") + seeded_chat = _chat([{"role": "assistant", "content": "same"}], "next", offset=1) + chat_histories = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[first_chat, seeded_chat]) + ).chat_completions_histories() + seeded_chat_source = chat_histories[1].message_sources[0] + assert seeded_chat_source is not None + assert seeded_chat_source.exchange is seeded_chat + assert seeded_chat_source.request_index == 0 + + first_message = _message() + seeded_message = _message() + seeded_message.start_time, seeded_message.end_time = _times(1) + seeded_message.request["messages"] = [{"role": "assistant", "content": "Hi"}] + message_histories = art.Trajectory( + exchanges=TrajectoryExchanges(messages=[first_message, seeded_message]) + ).anthropic_messages_histories() + seeded_message_source = message_histories[1].message_sources[0] + assert seeded_message_source is not None + assert seeded_message_source.exchange is seeded_message + assert seeded_message_source.request_index == 0 + + first_response = _response("response-1", "same") + seeded_response = _response("response-2", "next", offset=1) + seeded_response.request["input"] = [ + first_response.response.output[0].model_dump(mode="json", exclude_none=True) + ] + response_histories = art.Trajectory( + exchanges=TrajectoryExchanges(responses=[first_response, seeded_response]) + ).responses_histories() + seeded_response_source = response_histories[1].input_sources[0] + assert seeded_response_source is not None + assert seeded_response_source.exchange is seeded_response + assert seeded_response_source.request_index == 0 + + +def test_history_mutation_must_keep_source_sidecar_consistent() -> None: + trajectory = art.Trajectory( + exchanges=TrajectoryExchanges( + chat_completions=[_chat([{"role": "user", "content": "one"}], "first")] + ) + ) + history = trajectory.chat_completions_history() + history.messages.append({"role": "user", "content": "next"}) + with pytest.raises(ValueError, match="differ in length"): + history.tokenize() + + history.message_sources.append(None) + history.messages[0] = {"role": "user", "content": "edited"} + with pytest.raises(ValueError, match="no longer matches"): + history.tokenize() + + +def test_history_accepts_user_authored_messages_with_none_source() -> None: + trajectory = art.Trajectory( + exchanges=TrajectoryExchanges( + chat_completions=[_chat([{"role": "user", "content": "one"}], "first")] + ) + ) + history = trajectory.chat_completions_history() + history.messages.append({"role": "user", "content": "next"}) + history.message_sources.append(None) + + class Tokenizer: + def __call__(self, text: str, *, add_special_tokens: bool = False) -> list[int]: + assert not add_special_tokens + return {"one": [10], "first": [20], "next": [30]}[text] + + def apply_chat_template( + self, + messages: list[dict[str, Any]], + *, + tools: object, + tokenize: bool, + add_generation_prompt: bool, + chat_template: str | None = None, + **kwargs: object, + ) -> list[int]: + del tools, tokenize, add_generation_prompt, chat_template, kwargs + return [10] if len(messages) == 1 else [10, 20, 30] + + tokenized = history.tokenize(tokenizer=Tokenizer()) + + assert tokenized.token_ids == [10, 20, 30] + assert tokenized.flags[1] == art.TokenFlag.SAMPLED + + def test_protocol_histories_convert_to_chat_and_history_rejects_ambiguity() -> None: message_trajectory = art.Trajectory( exchanges=TrajectoryExchanges(messages=[_message()]) ) messages_history = message_trajectory.anthropic_messages_history() assert messages_history.system == "Be concise" - assert ( - art.AnthropicMessagesHistory.model_validate_json( - messages_history.model_dump_json() - ) - == messages_history - ) + assert not hasattr(messages_history, "model_dump") assert [message["role"] for message in messages_history.messages] == [ "user", "assistant", @@ -244,6 +576,48 @@ def test_protocol_histories_convert_to_chat_and_history_rejects_ambiguity() -> N assert isinstance(mixed.anthropic_messages_history(), art.AnthropicMessagesHistory) +def test_anthropic_chat_conversion_preserves_sources_for_expanded_messages() -> None: + exchange = _message() + exchange.request["messages"] = [ + { + "role": "user", + "content": [ + { + "type": "tool_result", + "tool_use_id": "call-1", + "content": "result", + }, + {"type": "text", "text": "continue"}, + ], + } + ] + history = art.Trajectory( + exchanges=TrajectoryExchanges(messages=[exchange]) + ).anthropic_messages_history() + + converted = history.as_chat_completions_history() + + assert [message["role"] for message in converted.messages] == [ + "system", + "tool", + "user", + "assistant", + ] + for source in converted.message_sources[1:3]: + assert source is not None + assert source.exchange is exchange + assert source.request_index == 0 + assert source.output_indices is None + response_source = converted.message_sources[-1] + assert response_source is not None + assert response_source.exchange is exchange + assert response_source.output_indices == (0,) + + converted.messages[2] = {"role": "user", "content": "changed"} + with pytest.raises(ValueError, match="no longer matches"): + converted.tokenize() + + def test_responses_history_expands_previous_response_chain() -> None: trajectory = art.Trajectory( exchanges=TrajectoryExchanges( @@ -261,9 +635,6 @@ def test_responses_history_expands_previous_response_chain() -> None: history = trajectory.responses_history() assert len(history.input) == 5 - assert ( - art.ResponsesHistory.model_validate_json(history.model_dump_json()) == history - ) chat_history = history.as_chat_completions_history() assert [message["role"] for message in chat_history.messages] == [ "user", @@ -272,17 +643,355 @@ def test_responses_history_expands_previous_response_chain() -> None: "assistant", ] assert dict(chat_history.messages[1]).get("reasoning") == "think" - assert ( - art.ChatCompletionsHistory.model_validate_json( - chat_history.model_dump_json(warnings="error") - ) - == chat_history - ) + assert all(source is not None for source in chat_history.message_sources) + assert chat_history.message_sources[1] is not None + assert chat_history.message_sources[1].output_indices == (0, 1) trajectory.exchanges.responses[1].request["previous_response_id"] = "missing" - with pytest.raises(ValueError, match="outside this history"): + external = trajectory.responses_histories() + assert len(external) == 2 + assert external[1].previous_response_id == "missing" + + +def test_responses_chat_conversion_preserves_request_and_output_sources() -> None: + exchange = _response("response-mixed-source", "answer") + exchange.request["input"] = [ + { + "id": "request-reasoning", + "type": "reasoning", + "summary": [{"type": "summary_text", "text": "prior thought"}], + } + ] + + converted = ( + art.Trajectory(exchanges=TrajectoryExchanges(responses=[exchange])) + .responses_history() + .as_chat_completions_history() + ) + + assert converted.messages == [ + { + "role": "assistant", + "content": "", + "reasoning": "prior thought", + }, + { + "role": "assistant", + "content": "answer", + }, + ] + request_source, output_source = converted.message_sources + assert request_source is not None + assert request_source.exchange is exchange + assert request_source.request_index == 0 + assert request_source.output_indices is None + assert output_source is not None + assert output_source.exchange is exchange + assert output_source.request_index is None + assert output_source.output_indices == (0,) + + +def test_responses_chat_conversion_splits_cross_exchange_assistant_sources() -> None: + first = _response("response-reasoning", "", reasoning="think") + first_data = first.response.model_dump(mode="python") + first_data["output"] = first_data["output"][:1] + first.response = Response.model_validate(first_data) + second = _response( + "response-answer", + "answer", + previous_response_id="response-reasoning", + offset=1, + ) + second.request["input"] = [] + + converted = ( + art.Trajectory(exchanges=TrajectoryExchanges(responses=[first, second])) + .responses_history() + .as_chat_completions_history() + ) + + assert converted.messages[-2:] == [ + {"role": "assistant", "content": "", "reasoning": "think"}, + {"role": "assistant", "content": "answer"}, + ] + first_source, second_source = converted.message_sources[-2:] + assert first_source is not None and first_source.exchange is first + assert first_source.output_indices == (0,) + assert second_source is not None and second_source.exchange is second + assert second_source.output_indices == (0,) + + +def test_responses_chat_conversion_owns_request_tool_group_by_first_item() -> None: + exchange = _response("response-request-tools", "answer") + exchange.request["input"] = [ + { + "type": "function_call", + "call_id": f"call-{index}", + "name": f"tool_{index}", + "arguments": "{}", + } + for index in range(2) + ] + + converted = ( + art.Trajectory(exchanges=TrajectoryExchanges(responses=[exchange])) + .responses_history() + .as_chat_completions_history() + ) + + assert len(converted.messages[0].get("tool_calls", [])) == 2 + source = converted.message_sources[0] + assert source is not None + assert source.exchange is exchange + assert source.request_index == 0 + assert source.output_indices is None + + +def test_responses_history_propagates_opaque_context_and_first_sources() -> None: + first = _response( + "response-1", + "first", + previous_response_id="outside-trajectory", + ) + first.request["conversation"] = "conversation-1" + second = _response( + "response-2", + "second", + previous_response_id="response-1", + offset=1, + ) + trajectory = art.Trajectory( + exchanges=TrajectoryExchanges(responses=[first, second]) + ) + + history = trajectory.responses_history() + + assert history.previous_response_id == "outside-trajectory" + assert history.conversation == "conversation-1" + assert history.input_sources[1] is not None + assert history.input_sources[1].exchange is first + assert history.input_sources[1].output_index == 0 + + +def test_branch_context_sources_follow_the_request_that_supplied_the_context() -> None: + first_message = _message() + second_message = _message() + second_message.start_time, second_message.end_time = _times(1) + second_message.request["system"] = "New instructions" + second_message.request["messages"] = [ + {"role": "user", "content": "Hello"}, + {"role": "assistant", "content": "Hi"}, + {"role": "user", "content": "Again"}, + ] + message_histories = art.Trajectory( + exchanges=TrajectoryExchanges(messages=[first_message, second_message]) + ).anthropic_messages_histories() + + assert message_histories[-1].system == "New instructions" + assert message_histories[-1].system_source is second_message + + first_response = _response("response-1", "first") + first_response.request["instructions"] = "Old instructions" + second_response = _response( + "response-2", + "second", + previous_response_id="response-1", + offset=1, + ) + second_response.request["instructions"] = "New instructions" + response_histories = art.Trajectory( + exchanges=TrajectoryExchanges(responses=[first_response, second_response]) + ).responses_histories() + + assert response_histories[-1].instructions == "New instructions" + assert response_histories[-1].instructions_source is second_response + + +def test_responses_history_maps_and_validates_generation_sources() -> None: + exchange = _response("response-1", "first", reasoning="think") + assert exchange.response.__pydantic_extra__ is not None + exchange.response.__pydantic_extra__["token_generations"] = [ + { + "prompt_token_ids": [1], + "output_tokens": [{"token_id": 2, "logprob": -0.1}], + "output_indices": [0, 1], + } + ] + + trajectory = art.Trajectory(exchanges=TrajectoryExchanges(responses=[exchange])) + history = trajectory.responses_history() + + assert { + source.generation_index + for source in history.input_sources + if source is not None and source.output_index is not None + } == {0} + + exchange.response.__pydantic_extra__["token_generations"][0]["output_indices"] = [0] + with pytest.raises(ValueError, match="every sampled output item"): trajectory.responses_history() + exchange.response.__pydantic_extra__["token_generations"] = [ + { + "prompt_token_ids": [1], + "output_tokens": [{"token_id": 2}], + "output_indices": [0], + }, + { + "prompt_token_ids": [1, 2], + "output_tokens": [{"token_id": 3}], + "output_indices": [0], + }, + ] + with pytest.raises(ValueError, match="nonoverlapping"): + trajectory.responses_history() + + +def test_cross_exchange_responses_reasoning_stripping_splits_histories() -> None: + first = _response("response-1", "first", reasoning="think") + first_data = first.response.model_dump(mode="python") + first_data["token_generations"] = [ + { + "prompt_token_ids": [1], + "output_tokens": [ + {"token_id": 2, "logprob": -0.2}, + {"token_id": 3, "logprob": -0.3}, + ], + "output_indices": [0, 1], + } + ] + first.response = Response.model_validate(first_data) + second = _response( + "response-2", + "second", + previous_response_id="response-1", + offset=1, + ) + second_data = second.response.model_dump(mode="python") + second_data["token_generations"] = [ + { + "prompt_token_ids": [1, 3, 4], + "output_tokens": [{"token_id": 5, "logprob": -0.5}], + "output_indices": [0], + } + ] + second.response = Response.model_validate(second_data) + + histories = art.Trajectory( + exchanges=TrajectoryExchanges(responses=[first, second]) + ).responses_histories() + + assert len(histories) == 2 + assert any(item.get("type") == "reasoning" for item in histories[0].input) + assert all(item.get("type") != "reasoning" for item in histories[1].input) + first_answer_source = histories[1].input_sources[1] + assert first_answer_source is not None + assert first_answer_source.exchange is first + assert first_answer_source.generation_index == 0 + + +@pytest.mark.parametrize( + ("second_prompt", "reasoning", "expected_histories"), + [([1, 2, 3], False, 1), ([1, 3], True, 2)], +) +def test_responses_multi_generation_history_leaves_match_tokenization( + second_prompt: list[int], reasoning: bool, expected_histories: int +) -> None: + exchange = _response( + "response-1", "first", reasoning="think" if reasoning else None + ) + data = exchange.response.model_dump(mode="python") + first_output_indices = list(range(len(data["output"]))) + tool_output_index = len(data["output"]) + second_output_index = tool_output_index + 1 + data["output"].extend( + [ + { + "id": "tool-output", + "type": "function_call_output", + "call_id": "call-1", + "output": "result", + "status": "completed", + }, + { + "id": "message-second", + "type": "message", + "role": "assistant", + "status": "completed", + "content": [ + { + "type": "output_text", + "text": "second", + "annotations": [], + "logprobs": [], + } + ], + }, + ] + ) + data["token_generations"] = [ + { + "prompt_token_ids": [1], + "output_tokens": [{"token_id": 2, "logprob": -0.2}], + "output_indices": first_output_indices, + }, + { + "prompt_token_ids": second_prompt, + "output_tokens": [{"token_id": 4, "logprob": -0.4}], + "output_indices": [second_output_index], + }, + ] + exchange.response = Response.model_validate(data) + trajectory = art.Trajectory(exchanges=TrajectoryExchanges(responses=[exchange])) + + histories = trajectory.responses_histories() + tokenized = trajectory.tokenize(multi_history=True) + + assert len(histories) == expected_histories + assert len(tokenized.histories) == expected_histories + final_sources = histories[-1].input_sources + assert final_sources[-2] is not None + assert final_sources[-2].generation_index is None + assert final_sources[-1] is not None + assert final_sources[-1].generation_index == 1 + if expected_histories == 2: + assert final_sources[1] is not None + assert final_sources[1].generation_index is None + assert all(item.get("type") != "reasoning" for item in histories[-1].input) + converted = histories[-1].as_chat_completions_history() + assert any( + source is not None + and source.generation_index == 1 + and source.output_indices == (second_output_index,) + for source in converted.message_sources + ) + + +def test_responses_generation_only_chat_source_has_empty_output_indices() -> None: + exchange = _response("response-empty-generation", "") + data = exchange.response.model_dump(mode="python") + data["output"] = [] + data["token_generations"] = [ + { + "prompt_token_ids": [1], + "output_tokens": [{"token_id": 2, "logprob": -0.2}], + "output_indices": [], + } + ] + exchange.response = Response.model_validate(data) + + converted = ( + art.Trajectory(exchanges=TrajectoryExchanges(responses=[exchange])) + .responses_history() + .as_chat_completions_history() + ) + + source = converted.message_sources[-1] + assert converted.messages[-1] == {"role": "assistant", "content": ""} + assert source is not None + assert source.output_indices == () + assert source.generation_index == 0 + def test_completions_history_preserves_exact_tokens_and_sampled_spans() -> None: trajectory = art.Trajectory( @@ -294,25 +1003,458 @@ def test_completions_history_preserves_exact_tokens_and_sampled_spans() -> None: ) ) - history = trajectory.completions_history() - assert history.token_ids == [1, 2, 3, 4] + history = trajectory.completions_token_history() + assert history.prompt == [1, 2, 3, 4] assert history.sampled_spans == [(1, 2), (3, 4)] with pytest.raises(ValueError, match="no chat-message structure"): history.as_chat_completions_history() -def test_completions_history_uses_request_token_ids_and_rejects_echo() -> None: +def test_completions_history_uses_request_token_ids() -> None: exchange = _completion([1], [2]) response = exchange.response.model_dump(mode="python") response["choices"][0].pop("prompt_token_ids") exchange.response = Completion.model_validate(response) trajectory = art.Trajectory(exchanges=TrajectoryExchanges(completions=[exchange])) - assert trajectory.completions_history().token_ids == [1, 2] + assert trajectory.completions_token_history().prompt == [1, 2] - exchange.request["echo"] = True - with pytest.raises(ValueError, match="echo=True"): - trajectory.completions_history() + +def test_batched_completions_create_every_prompt_choice_history() -> None: + exchange = _completion([1], [10]) + exchange.request["prompt"] = ["first", "second"] + response = exchange.response.model_dump(mode="python") + response["choices"] = [ + { + "index": index, + "finish_reason": "stop", + "text": f"answer-{index}", + "prompt_token_ids": [prompt_id], + "token_ids": [100 + index], + } + for index, prompt_id in enumerate((1, 1, 2, 2)) + ] + exchange.response = Completion.model_validate(response) + trajectory = art.Trajectory(exchanges=TrajectoryExchanges(completions=[exchange])) + + histories = trajectory.completions_token_histories() + + assert [history.prompt for history in histories] == [ + [1, 100], + [1, 101], + [2, 102], + [2, 103], + ] + with pytest.raises(ValueError, match="exactly one history"): + trajectory.history() + + +def test_batched_completions_prefer_exact_prompt_association() -> None: + exchange = _completion([1], [10]) + exchange.request["prompt"] = [[1], [2]] + response = exchange.response.model_dump(mode="python") + response["choices"] = [ + { + "index": 0, + "finish_reason": "stop", + "text": "second", + "prompt_token_ids": [2], + "token_ids": [20], + }, + { + "index": 1, + "finish_reason": "stop", + "text": "first", + "prompt_token_ids": [1], + "token_ids": [10], + }, + ] + exchange.response = Completion.model_validate(response) + trajectory = art.Trajectory(exchanges=TrajectoryExchanges(completions=[exchange])) + + histories = trajectory.completions_token_histories() + + assert [history.prompt for history in histories] == [[1, 10], [2, 20]] + + +def test_batched_completions_honor_interleaved_explicit_prompt_indices() -> None: + exchange = _completion([1], [10]) + exchange.request["prompt"] = ["first", "second"] + response = exchange.response.model_dump(mode="python") + response["choices"] = [ + { + "index": choice_index, + "finish_reason": "stop", + "text": text, + "prompt_index": prompt_index, + } + for choice_index, prompt_index, text in ( + (0, 1, "B0"), + (1, 0, "A0"), + (2, 1, "B1"), + (3, 0, "A1"), + ) + ] + exchange.response = Completion.model_validate(response) + + histories = art.Trajectory( + exchanges=TrajectoryExchanges(completions=[exchange]) + ).completions_string_histories() + + assert [history.prompt for history in histories] == [ + "firstA0", + "firstA1", + "secondB0", + "secondB1", + ] + + +@pytest.mark.parametrize(("prompt_index", "raises"), [(0, False), (1, True)]) +def test_batched_completions_validate_partial_prompt_index_fallback( + prompt_index: int, raises: bool +) -> None: + exchange = _completion([1], [10]) + exchange.request["prompt"] = ["first", "second"] + response = exchange.response.model_dump(mode="python") + response["choices"] = [ + { + "index": index, + "finish_reason": "stop", + "text": text, + **({"prompt_index": prompt_index} if index == 0 else {}), + } + for index, text in enumerate(("A0", "A1", "B0", "B1")) + ] + exchange.response = Completion.model_validate(response) + trajectory = art.Trajectory(exchanges=TrajectoryExchanges(completions=[exchange])) + + if raises: + with pytest.raises(ValueError, match="contradicts choice indices"): + trajectory.completions_string_histories() + return + assert [ + history.prompt for history in trajectory.completions_string_histories() + ] == ["firstA0", "firstA1", "secondB0", "secondB1"] + + +@pytest.mark.parametrize("prompt_index", [-1, 2, True, None]) +def test_batched_completions_reject_invalid_explicit_prompt_index( + prompt_index: object, +) -> None: + exchange = _completion([1], [10]) + exchange.request["prompt"] = ["first", "second"] + response = exchange.response.model_dump(mode="python") + response["choices"] = [ + { + "index": 0, + "finish_reason": "stop", + "text": "answer", + "prompt_index": prompt_index, + }, + { + "index": 1, + "finish_reason": "stop", + "text": "answer", + "prompt_index": 1, + }, + ] + exchange.response = Completion.model_validate(response) + + with pytest.raises(ValueError, match="prompt_index"): + art.Trajectory( + exchanges=TrajectoryExchanges(completions=[exchange]) + ).completions_string_histories() + + +def test_batched_completions_reject_prompt_index_exact_evidence_contradiction() -> None: + exchange = _completion([1], [10]) + exchange.request["prompt"] = [[1], [2]] + response = exchange.response.model_dump(mode="python") + response["choices"] = [ + { + "index": 0, + "finish_reason": "stop", + "text": "answer-0", + "prompt_index": 0, + "prompt_token_ids": [2], + "token_ids": [20], + }, + { + "index": 1, + "finish_reason": "stop", + "text": "answer-1", + "prompt_index": 1, + "prompt_token_ids": [1], + "token_ids": [10], + }, + ] + exchange.response = Completion.model_validate(response) + + with pytest.raises(ValueError, match="contradicts exact prompt evidence"): + art.Trajectory( + exchanges=TrajectoryExchanges(completions=[exchange]) + ).completions_token_histories() + + +def test_batched_completions_trust_explicit_string_prompt_index() -> None: + from art.trajectories._history import _completion_choice_groups + + exchange = _completion([1], [10]) + # Captured provider payloads can be broader than the SDK's prompt union. + exchange.request["prompt"] = cast(Any, ["same tokenization", [42]]) + response = exchange.response.model_dump(mode="python") + response["choices"] = [ + { + "index": index, + "finish_reason": "stop", + "text": f"answer-{index}", + "prompt_index": index, + "prompt_token_ids": [42], + "token_ids": [index + 10], + } + for index in range(2) + ] + exchange.response = Completion.model_validate(response) + + assert [ + [choice.index for choice in group] + for group in _completion_choice_groups(exchange) + ] == [[0], [1]] + + +def test_completions_histories_never_silently_omit_mixed_evidence() -> None: + exact = _completion([1], [2]) + missing = _completion([3], [4], offset=1) + missing.request["prompt"] = "question" + response = missing.response.model_dump(mode="python") + response["choices"][0].pop("prompt_token_ids") + response["choices"][0].pop("token_ids") + missing.response = Completion.model_validate(response) + trajectory = art.Trajectory( + exchanges=TrajectoryExchanges(completions=[exact, missing]) + ) + + with pytest.raises(ValueError, match="text prompts for every choice"): + trajectory.histories() + + +def test_completions_reject_ambiguous_batches_and_suffix() -> None: + ambiguous = _completion([1], [10]) + ambiguous.request["prompt"] = ["first", "second"] + with pytest.raises(ValueError, match="associate Completions choices"): + art.Trajectory( + exchanges=TrajectoryExchanges(completions=[ambiguous]) + ).histories() + + insertion = _completion([1], [10]) + insertion.request["suffix"] = "tail" + with pytest.raises(ValueError, match="suffix is not supported"): + art.Trajectory( + exchanges=TrajectoryExchanges(completions=[insertion]) + ).histories() + + +def test_completions_reject_duplicate_choice_indices() -> None: + exchange = _completion([1], [10]) + data = exchange.response.model_dump(mode="python") + data["choices"].append(dict(data["choices"][0])) + exchange.response = Completion.model_validate(data) + + with pytest.raises(ValueError, match="choice indices"): + art.Trajectory( + exchanges=TrajectoryExchanges(completions=[exchange]) + ).completions_token_histories() + + +def test_malformed_anthropic_content_raises_value_error() -> None: + exchange = _message() + exchange.request = cast( + MessagesRequest, + {"messages": [{"role": "user", "content": None}]}, + ) + + with pytest.raises(ValueError, match="Anthropic message content"): + art.Trajectory( + exchanges=TrajectoryExchanges(messages=[exchange]) + ).anthropic_messages_histories() + + +def test_tokenless_responses_generation_raises_value_error() -> None: + exchange = _response("response-1", "first") + data = exchange.response.model_dump(mode="python") + data["output"] = [ + { + "id": "tool-output", + "type": "function_call_output", + "call_id": "call-1", + "output": "result", + "status": "completed", + } + ] + data["token_generations"] = [ + { + "prompt_token_ids": [1], + "output_indices": [0], + } + ] + exchange.response = Response.model_validate(data) + + with pytest.raises(ValueError, match="without exact output tokens"): + art.Trajectory( + exchanges=TrajectoryExchanges(responses=[exchange]) + ).responses_histories() + + +def test_reasoning_stripping_produces_truthful_history_per_generation() -> None: + def exchange( + offset: int, + request_messages: list[dict[str, object]], + answer: str, + ) -> MessagesExchange: + start, end = _times(offset) + return MessagesExchange( + request={ + "model": "test/model", + "messages": request_messages, + "max_tokens": 16, + "thinking": {"type": "enabled", "budget_tokens": 8}, + }, + response=Message.model_validate( + { + "id": f"message-{offset}", + "type": "message", + "role": "assistant", + "model": "test/model", + "content": [ + { + "type": "thinking", + "thinking": f"thought-{offset}", + "signature": "sig", + }, + {"type": "text", "text": answer}, + ], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 1, "output_tokens": 1}, + } + ), + start_time=start, + end_time=end, + ) + + first = exchange(0, [{"role": "user", "content": "one"}], "first") + second = exchange( + 1, + [ + {"role": "user", "content": "one"}, + {"role": "assistant", "content": "first"}, + {"role": "user", "content": "two"}, + ], + "second", + ) + third = exchange( + 2, + [ + {"role": "user", "content": "one"}, + {"role": "assistant", "content": "first"}, + {"role": "user", "content": "two"}, + {"role": "assistant", "content": "second"}, + {"role": "user", "content": "three"}, + ], + "third", + ) + trajectory = art.Trajectory( + exchanges=TrajectoryExchanges(messages=[first, second, third]) + ) + + histories = trajectory.anthropic_messages_histories() + + assert len(histories) == 3 + assert [len(history.messages) for history in histories] == [2, 4, 6] + assert histories[1].message_sources[1] is not None + assert histories[1].message_sources[1].exchange is first + assert histories[1].message_sources[1].request_index is None + with pytest.raises(ValueError, match="exactly one history"): + trajectory.tokenize() + + +def test_chat_template_stripped_reasoning_splits_exact_histories() -> None: + first = _chat([{"role": "user", "content": "one"}], "first") + first_data = first.response.model_dump(mode="python") + first_data["prompt_token_ids"] = [1] + first_data["choices"][0]["message"]["reasoning"] = "thought-one" + first_data["choices"][0]["token_ids"] = [2, 3] + first.response = ChatCompletion.model_validate(first_data) + + second = _chat( + [ + {"role": "user", "content": "one"}, + {"role": "assistant", "content": "first"}, + {"role": "user", "content": "two"}, + ], + "second", + offset=1, + ) + second_data = second.response.model_dump(mode="python") + second_data["prompt_token_ids"] = [1, 3, 4] + second_data["choices"][0]["message"]["reasoning"] = "thought-two" + second_data["choices"][0]["token_ids"] = [5, 6] + second.response = ChatCompletion.model_validate(second_data) + trajectory = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[first, second]) + ) + + histories = trajectory.chat_completions_histories() + + assert len(histories) == 2 + assert [len(history.messages) for history in histories] == [2, 4] + assert histories[1].message_sources[1] is not None + assert histories[1].message_sources[1].exchange is first + assert histories[1].message_sources[1].choice_index == 0 + with pytest.raises(ValueError, match="exactly one history"): + trajectory.tokenize() + + +def test_reasoning_stripped_tool_call_keeps_first_sampled_source() -> None: + first = _chat([{"role": "user", "content": "one"}], "") + first_data = first.response.model_dump(mode="python") + first_data["choices"][0]["message"] = { + "role": "assistant", + "content": None, + "reasoning": "thought", + "tool_calls": [ + { + "id": "call-1", + "type": "function", + "function": {"name": "lookup", "arguments": "{}"}, + } + ], + } + first.response = ChatCompletion.model_validate(first_data) + second = _chat( + [ + {"role": "user", "content": "one"}, + { + "role": "assistant", + "content": None, + "tool_calls": first_data["choices"][0]["message"]["tool_calls"], + }, + {"role": "user", "content": "two"}, + ], + "second", + offset=1, + ) + + histories = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[first, second]) + ).chat_completions_histories() + + assert len(histories) == 2 + source = histories[1].message_sources[1] + assert source is not None + assert source.exchange is first + assert source.choice_index == 0 + assert source.request_index is None def test_history_rejects_mutated_mixed_representation() -> None: @@ -332,22 +1474,27 @@ def test_legacy_messages_delegate_through_history() -> None: messages_and_choices=[{"role": "user", "content": "hello"}] ) - assert isinstance(trajectory.history(), art.History) + assert isinstance(trajectory.history(), art.LegacyHistory) assert trajectory.messages() == [{"role": "user", "content": "hello"}] - with pytest.raises(ValueError, match="do not identify a model"): - trajectory.history(model="test/model") + assert isinstance(trajectory.history(model="test/model"), art.LegacyHistory) + with pytest.raises(ValueError, match="requires model="): + trajectory.tokenize() def test_legacy_messages_preserve_primary_history_with_additional_histories() -> None: trajectory = art.Trajectory( messages_and_choices=[{"role": "user", "content": "primary"}], additional_histories=[ - art.History(messages_and_choices=[{"role": "user", "content": "alternate"}]) + art.LegacyHistory( + messages_and_choices=[{"role": "user", "content": "alternate"}] + ) ], ) assert trajectory.messages() == [{"role": "user", "content": "primary"}] - with pytest.raises(ValueError, match="multiple legacy histories"): + assert len(trajectory.histories()) == 2 + assert len(trajectory.chat_completions_histories()) == 2 + with pytest.raises(ValueError, match="exactly one history"): trajectory.history() diff --git a/tests/unit/trajectories/test_tokenize.py b/tests/unit/trajectories/test_tokenize.py index 7dc2d80d8..914eadaa5 100644 --- a/tests/unit/trajectories/test_tokenize.py +++ b/tests/unit/trajectories/test_tokenize.py @@ -3,25 +3,37 @@ import builtins from datetime import datetime, timedelta import math -from types import SimpleNamespace +import random +import re +from statistics import median +import sys +from time import perf_counter +from types import ModuleType, SimpleNamespace from typing import Any from anthropic.types import ImageBlockParam, Message, MessageParam from openai.types import Completion -from openai.types.chat import ChatCompletion +from openai.types.chat import ChatCompletion, ChatCompletionMessageParam from openai.types.chat.chat_completion_token_logprob import ChatCompletionTokenLogprob -from openai.types.responses import Response +from openai.types.responses import ( + EasyInputMessageParam, + Response, + ResponseInputParam, + ResponseOutputMessageParam, +) import pytest import art from art.trajectories import ( ChatCompletionsExchange, + ChatCompletionsMessageSource, ChatCompletionsRequest, CompletionsExchange, CompletionsRequest, MessagesExchange, MessagesRequest, ResponsesExchange, + ResponsesItemSource, ResponsesRequest, TrajectoryExchanges, ) @@ -34,6 +46,11 @@ def _chat_exchange( model: str = "test/model", offset: int = 0, ) -> ChatCompletionsExchange: + messages: list[ChatCompletionMessageParam] = [] + for turn in range(offset + 1): + messages.append({"role": "user", "content": f"turn {turn}"}) + if turn < offset: + messages.append({"role": "assistant", "content": "answer"}) response = ChatCompletion.model_validate( { "id": f"chat-{offset}", @@ -66,7 +83,7 @@ def _chat_exchange( return ChatCompletionsExchange( request=ChatCompletionsRequest( model=model, - messages=[{"role": "user", "content": f"turn {offset}"}], + messages=messages, ), response=response, start_time=start, @@ -89,7 +106,7 @@ def _completion_exchange( { "index": 0, "finish_reason": "stop", - "text": "answer", + "text": f"{'question' if echo else ''}answer", "prompt_token_ids": [1], "token_ids": [2], "logprobs": { @@ -113,6 +130,38 @@ def _completion_exchange( ) +def _message_exchange( + request: MessagesRequest, + *, + identifier: str = "message-1", + content: list[dict[str, object]] | None = None, + duration: timedelta = timedelta(milliseconds=1), + offset: int = 0, + response_model: str = "test/model", + **response_extra: object, +) -> MessagesExchange: + start = datetime(2026, 1, 1) + timedelta(seconds=offset) + response = Message.model_validate( + { + "id": identifier, + "type": "message", + "role": "assistant", + "model": response_model, + "content": content or [{"type": "text", "text": "answer"}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 0, "output_tokens": 0}, + **response_extra, + } + ) + return MessagesExchange( + request=request, + response=response, + start_time=start, + end_time=start + duration, + ) + + def test_exact_tokens_form_one_append_only_history_without_tokenizer( monkeypatch: pytest.MonkeyPatch, ) -> None: @@ -138,7 +187,7 @@ def import_without_tokenizer_dependencies(name: str, *args: Any, **kwargs: Any): monkeypatch.delenv("WANDB_API_KEY", raising=False) monkeypatch.setattr(builtins, "__import__", import_without_tokenizer_dependencies) - tokenized = art.tokenize_trajectory(trajectory) + tokenized = trajectory.tokenize() assert tokenized.token_ids == [1, 2, 3, 4] assert tokenized.flags == [ @@ -153,6 +202,206 @@ def import_without_tokenizer_dependencies(name: str, *args: Any, **kwargs: Any): assert tokenized.logprobs[3] == -0.4 +def test_empty_tool_calls_normalization_preserves_exact_continuation() -> None: + first = _chat_exchange([1], [2]) + first_data = first.response.model_dump(mode="python") + first_data["choices"][0]["message"]["tool_calls"] = [] + first.response = ChatCompletion.model_validate(first_data) + second = _chat_exchange([1, 2, 3], [4], offset=1) + + tokenized = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[first, second]) + ).tokenize() + + assert tokenized.token_ids == [1, 2, 3, 4] + assert tokenized.logprobs[1::2] == [-0.2, -0.4] + + +def test_messages_exact_prompt_and_output_do_not_load_a_tokenizer( + monkeypatch: pytest.MonkeyPatch, +) -> None: + trajectory = art.Trajectory( + exchanges=TrajectoryExchanges( + messages=[ + _message_exchange( + MessagesRequest( + model="test/model", + messages=[{"role": "user", "content": "question"}], + max_tokens=16, + ), + identifier="message-exact", + prompt_token_ids=[1, 2], + token_ids=[3], + logprobs=[-0.3], + ) + ] + ) + ) + monkeypatch.setattr( + "art.trajectories._tokenize._load_tokenizer", + lambda _config: pytest.fail("exact Messages evidence loaded a tokenizer"), + ) + + tokenized = trajectory.tokenize() + + assert tokenized.token_ids == [1, 2, 3] + assert all(math.isnan(value) for value in tokenized.logprobs[:2]) + assert tokenized.logprobs[2] == -0.3 + assert tokenized.flags == [ + art.TokenFlag.EXACT, + art.TokenFlag.EXACT, + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + ] + + +def test_anthropic_cache_control_change_starts_new_source_lineage() -> None: + def exchange( + *, + identifier: str, + messages: list[MessageParam], + prompt_token_ids: list[int], + output_token_id: int, + offset: int, + ) -> MessagesExchange: + return _message_exchange( + MessagesRequest( + model="test/model", + messages=messages, + max_tokens=16, + ), + identifier=identifier, + content=[{"type": "text", "text": f"answer {offset}"}], + offset=offset, + prompt_token_ids=prompt_token_ids, + token_ids=[output_token_id], + logprobs=[-0.1], + ) + + first = exchange( + identifier="message-1", + messages=[{"role": "user", "content": [{"type": "text", "text": "question"}]}], + prompt_token_ids=[1], + output_token_id=2, + offset=0, + ) + second = exchange( + identifier="message-2", + messages=[ + { + "role": "user", + "content": [ + { + "type": "text", + "text": "question", + "cache_control": {"type": "ephemeral"}, + } + ], + }, + { + "role": "assistant", + "content": [{"type": "text", "text": "answer 0"}], + }, + {"role": "user", "content": "follow up"}, + ], + prompt_token_ids=[1, 2, 3], + output_token_id=4, + offset=1, + ) + + histories = art.Trajectory( + exchanges=TrajectoryExchanges(messages=[first, second]) + ).anthropic_messages_histories() + + assert len(histories) == 2 + updated = histories[1] + source = updated.message_sources[0] + assert source is not None + assert source.exchange is second + assert source.request_index == 0 + assert updated.tokenize().token_ids == [1, 2, 3, 4] + + +def test_converted_anthropic_system_history_preserves_exact_assistant_evidence( + monkeypatch: pytest.MonkeyPatch, +) -> None: + exchange = _message_exchange( + MessagesRequest( + model="test/model", + system="system", + messages=[{"role": "user", "content": "question"}], + max_tokens=16, + ), + identifier="message-system-exact", + prompt_token_ids=[10, 11], + token_ids=[12], + logprobs=[-0.12], + ) + trajectory = art.Trajectory(exchanges=TrajectoryExchanges(messages=[exchange])) + monkeypatch.setattr( + "art.trajectories._tokenize._load_tokenizer", + lambda _config: pytest.fail("exact converted history loaded a tokenizer"), + ) + + history = trajectory.anthropic_messages_history().as_chat_completions_history() + assert all( + source is None or source.exchange is exchange + for source in history.message_sources + ) + tokenized = history.tokenize() + + assert tokenized.token_ids == [10, 11, 12] + assert tokenized.logprobs[-1] == pytest.approx(-0.12) + assert tokenized.flags[-1] == art.TokenFlag.EXACT | art.TokenFlag.SAMPLED + + +def test_converted_anthropic_system_source_rejects_sampled_response_mutation() -> None: + exchange = _message_exchange( + MessagesRequest( + model="test/model", + system="system", + messages=[{"role": "user", "content": "question"}], + max_tokens=16, + ), + identifier="message-system-source", + prompt_token_ids=[10, 11], + token_ids=[12], + logprobs=[-0.12], + ) + history = ( + art.Trajectory(exchanges=TrajectoryExchanges(messages=[exchange])) + .anthropic_messages_history() + .as_chat_completions_history() + ) + history.messages[0] = {"role": "assistant", "content": "answer"} + + with pytest.raises(ValueError, match="no longer matches its source exchange"): + history.tokenize() + + +def test_converted_responses_history_tokenizes_without_native_chat_exchange( + monkeypatch: pytest.MonkeyPatch, +) -> None: + exchange = _response_exchange( + "response-converted-exact", 12, prompt_token_ids=[10, 11] + ) + trajectory = art.Trajectory(exchanges=TrajectoryExchanges(responses=[exchange])) + monkeypatch.setattr( + "art.trajectories._tokenize._load_tokenizer", + lambda _config: pytest.fail("exact converted history loaded a tokenizer"), + ) + + history = trajectory.responses_history().as_chat_completions_history() + assert all( + source is None or source.exchange is exchange + for source in history.message_sources + ) + tokenized = history.tokenize() + + assert tokenized.token_ids == [10, 11, 12] + assert tokenized.logprobs[-1] == pytest.approx(-0.1) + assert tokenized.flags[-1] == art.TokenFlag.EXACT | art.TokenFlag.SAMPLED + + def test_malformed_explicit_exact_token_metadata_fails_closed() -> None: chat = _chat_exchange([1], [2]) chat_extra = chat.response.choices[0].model_extra @@ -167,31 +416,16 @@ def test_malformed_explicit_exact_token_metadata_fails_closed() -> None: response = _response_exchange("response-invalid", 2) response_extra = response.response.model_extra assert response_extra is not None - response_extra["raw_output_tokens"] = [{"token_id": "invalid"}] + response_extra["token_generations"][0]["output_tokens"] = [{"token_id": "invalid"}] - message_response = Message.model_validate( - { - "id": "message-invalid", - "type": "message", - "role": "assistant", - "model": "test/model", - "content": [{"type": "text", "text": "answer"}], - "stop_reason": "end_turn", - "stop_sequence": None, - "usage": {"input_tokens": 1, "output_tokens": 1}, - "token_ids": [2, "invalid"], - } - ) - start = datetime(2026, 1, 1) - message = MessagesExchange( - request=MessagesRequest( + message = _message_exchange( + MessagesRequest( model="test/model", messages=[{"role": "user", "content": "question"}], max_tokens=16, ), - response=message_response, - start_time=start, - end_time=start + timedelta(milliseconds=1), + identifier="message-invalid", + token_ids=[2, "invalid"], ) trajectories = [ @@ -202,7 +436,20 @@ def test_malformed_explicit_exact_token_metadata_fails_closed() -> None: ] for trajectory in trajectories: with pytest.raises(ValueError, match="exact token"): - art.tokenize_trajectory(trajectory, base_model="base/model") + trajectory.tokenize(base_model="base/model") + + +@pytest.mark.parametrize("token_id", [-1, True]) +def test_exact_token_metadata_rejects_negative_ids_and_booleans( + token_id: object, +) -> None: + exchange = _response_exchange("response-invalid-token", 2) + extra = exchange.response.model_extra + assert extra is not None + extra["token_generations"][0]["output_tokens"] = [{"token_id": token_id}] + + with pytest.raises(ValueError, match="invalid exact token ID"): + art.Trajectory(exchanges=TrajectoryExchanges(responses=[exchange])).tokenize() @pytest.mark.parametrize( @@ -213,794 +460,4888 @@ def test_malformed_explicit_exact_token_metadata_fails_closed() -> None: _completion_exchange(echo=True), ], ) -def test_completions_reject_batch_prompts_and_echo( +def test_completions_support_single_item_batches_and_echo( exchange: CompletionsExchange, ) -> None: - with pytest.raises(ValueError, match="batched Completions|echo=True"): - art.tokenize_trajectory( - art.Trajectory(exchanges=TrajectoryExchanges(completions=[exchange])) - ) + tokenized = art.Trajectory( + exchanges=TrajectoryExchanges(completions=[exchange]) + ).tokenize() + assert tokenized.token_ids == [1, 2] -def test_branching_and_multiple_models_require_explicit_resolution() -> None: - branching = art.Trajectory( - exchanges=TrajectoryExchanges( - chat_completions=[ - _chat_exchange([1], [2], offset=0), - _chat_exchange([9], [3], offset=1), - ] - ) - ) - with pytest.raises(ValueError, match="append-only"): - art.tokenize_trajectory(branching) +def test_completions_echo_preserves_prompt_logprobs_without_sampling_them() -> None: + exchange = _completion_exchange(echo=True) + payload = exchange.response.model_dump(mode="python") + payload["choices"][0]["token_ids"] = [2] + payload["choices"][0]["logprobs"] = { + "tokens": ["token_id:1", "token_id:2"], + "token_logprobs": [-0.1, -0.2], + "top_logprobs": [{}, {}], + "text_offset": [0, 8], + } + exchange.response = Completion.model_validate(payload) - mixed = art.Trajectory( - exchanges=TrajectoryExchanges( - chat_completions=[ - _chat_exchange([1], [2], model="one", offset=0), - _chat_exchange([3], [4], model="two", offset=1), - ] - ) - ) - with pytest.raises(ValueError, match="exactly one model"): - art.tokenize_trajectory(mixed) - assert art.tokenize_trajectory(mixed, model="two").token_ids == [3, 4] + tokenized = art.Trajectory( + exchanges=TrajectoryExchanges(completions=[exchange]) + ).tokenize() + assert tokenized.token_ids == [1, 2] + assert tokenized.logprobs == [-0.1, -0.2] + assert tokenized.flags == [ + art.TokenFlag.EXACT, + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + ] -class _FakeTokenizer: - def __init__(self) -> None: - self.calls: list[dict[str, object]] = [] - def apply_chat_template( - self, messages: list[dict[str, Any]], **kwargs: object - ) -> list[int]: - self.calls.append(kwargs) - return [10, 11] if messages[-1]["role"] == "assistant" else [10] +def test_completions_echo_does_not_strip_repeated_prompt_token_from_completion() -> ( + None +): + exchange = _completion_exchange(echo=True) + payload = exchange.response.model_dump(mode="python") + payload["choices"][0]["text"] = "questionquestionanswer" + payload["choices"][0]["token_ids"] = [1, 2] + payload["choices"][0]["logprobs"] = { + "tokens": ["token_id:1", "token_id:2"], + "token_logprobs": [-0.2, -0.3], + "top_logprobs": [{}, {}], + "text_offset": [8, 16], + } + exchange.response = Completion.model_validate(payload) + class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> list[int]: + del kwargs + return {"question": [1], "questionanswer": [1, 2]}[text] -def test_fallback_uses_template_overrides_and_nan_logprobs( - monkeypatch: pytest.MonkeyPatch, -) -> None: - response = Message.model_validate( - { - "id": "msg_1", - "type": "message", - "role": "assistant", - "model": "test/model", - "content": [{"type": "text", "text": "answer"}], - "stop_reason": "end_turn", - "stop_sequence": None, - "usage": {"input_tokens": 1, "output_tokens": 1}, - } - ) - start = datetime(2026, 1, 1) - exchange = MessagesExchange( - request=MessagesRequest( - model="wandb-artifact:///entity/project/run:step0", - messages=[{"role": "user", "content": "question"}], - chat_template="request-template", - chat_template_kwargs={"request": True}, - thinking={"type": "enabled", "budget_tokens": 128}, - ), - response=response, - start_time=start, - end_time=start + timedelta(seconds=1), - ) - tokenizer = _FakeTokenizer() - loaded_base_models: list[str] = [] - monkeypatch.setattr( - "art.trajectories._tokenize._load_tokenizer", - lambda config: loaded_base_models.append(config.base_model) or tokenizer, - ) - monkeypatch.setattr( - "art.trajectories._tokenize._artifact_config", - lambda _model: pytest.fail("explicit base_model should bypass W&B"), - ) + def apply_chat_template( + self, messages: list[dict[str, object]], **kwargs: object + ) -> list[int]: + raise AssertionError( + "Completions tokenization must not render chat messages" + ) - result = art.tokenize_trajectory( - art.Trajectory(exchanges=TrajectoryExchanges(messages=[exchange])), - base_model="base/model", - chat_template="explicit-template", - chat_template_kwargs={"explicit": True}, - ) + tokenized = art.Trajectory( + exchanges=TrajectoryExchanges(completions=[exchange]) + ).tokenize(tokenizer=Tokenizer()) - assert result.token_ids == [10, 11] - assert loaded_base_models == ["base/model"] - assert result.flags == [art.TokenFlag(0), art.TokenFlag.SAMPLED] - assert math.isnan(result.logprobs[1]) - assert tokenizer.calls == [ - { - "tools": None, - "tokenize": True, - "add_generation_prompt": True, - "chat_template": "explicit-template", - "request": True, - "explicit": True, - "enable_thinking": True, - "thinking_budget": 128, - }, - { - "tools": None, - "tokenize": True, - "add_generation_prompt": False, - "chat_template": "explicit-template", - "request": True, - "explicit": True, - "enable_thinking": True, - "thinking_budget": 128, - }, + assert tokenized.token_ids == [1, 1, 2] + assert tokenized.flags == [ + art.TokenFlag.EXACT, + art.TokenFlag.SAMPLED, + art.TokenFlag.SAMPLED, ] + assert all(math.isnan(logprob) for logprob in tokenized.logprobs) -@pytest.mark.parametrize( - ("model", "artifact_name"), - [ - ("wandb-artifact:///entity/project/run", "entity/project/run:latest"), - ("wandb-artifact:///entity/project/run:step0", "entity/project/run:step0"), - ], -) -def test_checkpoint_fallback_preserves_artifact_version_and_renderer( - monkeypatch: pytest.MonkeyPatch, - model: str, - artifact_name: str, -) -> None: - artifact_names: list[str] = [] +def test_completions_echo_strips_prompt_from_proven_combined_token_carrier() -> None: + exchange = _completion_exchange(echo=True) + payload = exchange.response.model_dump(mode="python") + payload["choices"][0]["text"] = "questionquestionanswer" + payload["choices"][0]["token_ids"] = [1, 1, 2] + payload["choices"][0]["logprobs"] = { + "tokens": ["token_id:1", "token_id:2"], + "token_logprobs": [-0.2, -0.3], + "top_logprobs": [{}, {}], + "text_offset": [8, 16], + } + exchange.response = Completion.model_validate(payload) - class Api: - def artifact(self, name: str) -> SimpleNamespace: - artifact_names.append(name) - return SimpleNamespace( - metadata={ - "wandb.base_model": "base/model", - "renderer": { - "tokenizer_revision": "revision", - "chat_template": "template", - "chat_template_kwargs": {"thinking": True}, - }, - } - ) + tokenized = art.Trajectory( + exchanges=TrajectoryExchanges(completions=[exchange]) + ).tokenize() - monkeypatch.setattr("wandb.apis.public.Api", Api) - exchange = _chat_exchange([], [], model=model) - extra = exchange.response.choices[0].model_extra - assert extra is not None - extra.pop("prompt_token_ids") - extra.pop("token_ids") - exchange.response.choices[0].logprobs = None - tokenizer = _FakeTokenizer() - configs = [] - monkeypatch.setattr( - "art.trajectories._tokenize._load_tokenizer", - lambda config: configs.append(config) or tokenizer, - ) + assert tokenized.token_ids == [1, 1, 2] + assert tokenized.logprobs[1:] == pytest.approx([-0.2, -0.3]) + assert tokenized.flags == [ + art.TokenFlag.EXACT, + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + ] - art.tokenize_trajectory( - art.Trajectory(exchanges=TrajectoryExchanges(chat_completions=[exchange])) - ) - config = configs[0] - assert artifact_names == [artifact_name] - assert config.base_model == "base/model" - assert config.revision == "revision" - assert config.chat_template == "template" - assert config.chat_template_kwargs == {"thinking": True} - assert tokenizer.calls[0]["chat_template"] == "template" - assert tokenizer.calls[0]["thinking"] is True +def test_completions_echo_strips_prompt_from_proven_textual_carrier() -> None: + exchange = _completion_exchange(echo=True) + payload = exchange.response.model_dump(mode="python") + payload["choices"][0]["token_ids"] = [1, 2] + payload["choices"][0]["logprobs"] = { + "tokens": ["question", "answer"], + "token_logprobs": [-0.1, -0.2], + "top_logprobs": [{}, {}], + "text_offset": [0, 8], + } + exchange.response = Completion.model_validate(payload) + tokenized = art.Trajectory( + exchanges=TrajectoryExchanges(completions=[exchange]) + ).tokenize() -def test_anthropic_fallback_preserves_thinking_and_tool_history() -> None: + assert tokenized.token_ids == [1, 2] + assert tokenized.logprobs == pytest.approx([-0.1, -0.2]) + assert tokenized.flags == [ + art.TokenFlag.EXACT, + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + ] + + +def test_completions_echo_prefers_full_logprob_carrier_to_id_prefix_heuristic() -> None: + exchange = _completion_exchange(echo=True) + payload = exchange.response.model_dump(mode="python") + payload["choices"][0]["token_ids"] = [1, 2] + payload["choices"][0]["logprobs"] = { + "tokens": ["token_id:1", "token_id:1", "token_id:2"], + "token_logprobs": [-0.1, -0.2, -0.3], + "top_logprobs": [{}, {}, {}], + "text_offset": [0, 8, 9], + } + exchange.response = Completion.model_validate(payload) + + tokenized = art.Trajectory( + exchanges=TrajectoryExchanges(completions=[exchange]) + ).tokenize() + + assert tokenized.token_ids == [1, 1, 2] + assert tokenized.logprobs == [-0.1, -0.2, -0.3] + assert tokenized.flags == [ + art.TokenFlag.EXACT, + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + ] + + +def test_completions_echo_without_prompt_ids_falls_back_without_sampling_prompt() -> ( + None +): + exchange = _completion_exchange(echo=True) + payload = exchange.response.model_dump(mode="python") + payload["choices"][0].pop("prompt_token_ids") + payload["choices"][0].pop("token_ids") + payload["choices"][0]["logprobs"] = { + "tokens": ["token_id:1", "token_id:2"], + "token_logprobs": [-0.1, -0.2], + "top_logprobs": [{}, {}], + "text_offset": [0, 8], + } + exchange.response = Completion.model_validate(payload) + + class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> list[int]: + del kwargs + return {"question": [1], "answer": [2]}[text] + + def apply_chat_template( + self, messages: list[dict[str, Any]], **kwargs: object + ) -> list[int]: + raise AssertionError((messages, kwargs)) + + tokenized = art.Trajectory( + exchanges=TrajectoryExchanges(completions=[exchange]) + ).tokenize(tokenizer=Tokenizer()) + + assert tokenized.token_ids == [1, 2] + assert tokenized.flags == [art.TokenFlag(0), art.TokenFlag.SAMPLED] + assert all(math.isnan(logprob) for logprob in tokenized.logprobs) + + +def test_batched_completions_echo_uses_selected_prompt_boundary() -> None: + exchange = _completion_exchange(prompt=["p0", "p1"], echo=True) + payload = exchange.response.model_dump(mode="python") + payload["choices"] = [ + { + "index": 0, + "finish_reason": "stop", + "text": "p0a", + "logprobs": { + "tokens": ["p0", "a"], + "token_logprobs": [-9.0, -0.1], + "top_logprobs": [{}, {}], + "text_offset": [0, 2], + }, + }, + { + "index": 1, + "finish_reason": "stop", + "text": "p1b", + "logprobs": { + "tokens": ["p1", "b"], + "token_logprobs": [-9.0, -0.2], + "top_logprobs": [{}, {}], + "text_offset": [0, 2], + }, + }, + ] + exchange.request["n"] = 1 + exchange.response = Completion.model_validate(payload) + + class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> list[int]: + del kwargs + return {"p0": [1], "p1": [2], "a": [3], "b": [4]}[text] + + def apply_chat_template( + self, messages: list[dict[str, Any]], **kwargs: object + ) -> list[int]: + raise AssertionError((messages, kwargs)) + + histories = art.Trajectory( + exchanges=TrajectoryExchanges(completions=[exchange]) + ).completions_string_histories() + + assert [ + history.tokenize(tokenizer=Tokenizer()).token_ids for history in histories + ] == [ + [1, 3], + [2, 4], + ] + assert [ + history.tokenize(tokenizer=Tokenizer()).logprobs[-1] for history in histories + ] == pytest.approx([-0.1, -0.2]) + + +def test_completions_exact_ids_accept_textual_logprobs() -> None: + exchange = _completion_exchange() + payload = exchange.response.model_dump(mode="python") + payload["choices"][0]["logprobs"]["tokens"] = ["answer"] + payload["choices"][0]["logprobs"]["token_logprobs"] = [-0.75] + exchange.response = Completion.model_validate(payload) + + tokenized = art.Trajectory( + exchanges=TrajectoryExchanges(completions=[exchange]) + ).tokenize() + + assert tokenized.token_ids == [1, 2] + assert tokenized.logprobs[1] == -0.75 + + +def test_mutated_completions_token_prompt_drops_stale_exact_evidence() -> None: + history = art.Trajectory( + exchanges=TrajectoryExchanges(completions=[_completion_exchange()]) + ).completions_token_history() + history.prompt[0] = 99 + + tokenized = history.tokenize() + + assert tokenized.token_ids == [99, 2] + assert tokenized.flags == [ + art.TokenFlag(0), + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + ] + + +def test_mutated_completions_token_choice_is_rejected() -> None: + history = art.Trajectory( + exchanges=TrajectoryExchanges(completions=[_completion_exchange()]) + ).completions_token_history() + history.prompt[-1] = 99 + + with pytest.raises(ValueError, match="sampled output"): + history.tokenize() + + +def test_completions_history_rejects_model_and_sampled_span_mutation() -> None: + history = art.Trajectory( + exchanges=TrajectoryExchanges(completions=[_completion_exchange()]) + ).completions_token_history() + history.model = "other/model" + with pytest.raises(ValueError, match="model no longer matches"): + history.tokenize() + + history.model = "test/model" + history.sampled_spans = [(0, len(history.prompt))] + with pytest.raises(ValueError, match="exactly match choice-backed"): + history.tokenize() + + +@pytest.mark.parametrize( + "protocol", + ["chat_completions", "messages", "responses", "completions"], +) +def test_source_backed_histories_reject_model_mutation(protocol: str) -> None: + if protocol == "chat_completions": + history = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[_chat_exchange([1], [2])]) + ).chat_completions_history() + elif protocol == "messages": + history = art.Trajectory( + exchanges=TrajectoryExchanges( + messages=[ + _message_exchange( + MessagesRequest( + model="test/model", + messages=[{"role": "user", "content": "question"}], + max_tokens=16, + ), + prompt_token_ids=[1], + token_ids=[2], + logprobs=[-0.2], + ) + ] + ) + ).anthropic_messages_history() + elif protocol == "responses": + history = art.Trajectory( + exchanges=TrajectoryExchanges( + responses=[_response_exchange("response-model", 2)] + ) + ).responses_history() + else: + history = art.Trajectory( + exchanges=TrajectoryExchanges(completions=[_completion_exchange()]) + ).completions_token_history() + + history.model = "other/model" + + with pytest.raises(ValueError, match="model no longer matches"): + history.tokenize() + + +def test_completions_string_history_preserves_textual_logprobs() -> None: + exchange = _completion_exchange() + payload = exchange.response.model_dump(mode="python") + payload["choices"][0].pop("prompt_token_ids", None) + payload["choices"][0].pop("token_ids", None) + payload["choices"][0]["logprobs"]["tokens"] = ["answer"] + payload["choices"][0]["logprobs"]["token_logprobs"] = [-0.8] + exchange.response = Completion.model_validate(payload) + history = art.Trajectory( + exchanges=TrajectoryExchanges(completions=[exchange]) + ).completions_string_history() + + class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> list[int]: + del kwargs + return {"question": [1], "answer": [2]}[text] + + def apply_chat_template(self, *args: object, **kwargs: object) -> list[int]: + raise AssertionError("Completions tokenization does not render chat") + + tokenized = history.tokenize(tokenizer=Tokenizer()) + + assert tokenized.token_ids == [1, 2] + assert tokenized.logprobs[1] == -0.8 + assert tokenized.flags[1] == art.TokenFlag.SAMPLED + + +def test_mutated_completions_string_prompt_retokens_without_stale_exact() -> None: + history = art.Trajectory( + exchanges=TrajectoryExchanges(completions=[_completion_exchange()]) + ).completions_string_history() + history.prompt = "changed" + history.prompt[len("question") :] + first = history.prompt_sources[0] + history.prompt_sources[0] = type(first)( + start=0, + end=len("changed"), + source=first.source, + ) + shift = len("changed") - len("question") + second = history.prompt_sources[1] + history.prompt_sources[1] = type(second)( + start=second.start + shift, + end=second.end + shift, + source=second.source, + ) + history.sampled_spans = [ + (start + shift, end + shift) for start, end in history.sampled_spans + ] + + class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> list[int]: + del kwargs + return {"changed": [9], "answer": [2]}[text] + + def apply_chat_template(self, *args: object, **kwargs: object) -> list[int]: + raise AssertionError("Completions tokenization does not render chat") + + tokenized = history.tokenize(tokenizer=Tokenizer()) + + assert tokenized.token_ids == [9, 2] + assert tokenized.flags[0] == art.TokenFlag(0) + + +def test_batched_completions_map_each_choice_to_its_prompt() -> None: + exchange = _completion_exchange(prompt=["first", "second"]) + payload = exchange.response.model_dump(mode="python") + payload["choices"] = [ + { + **payload["choices"][0], + "index": 0, + "text": "one", + "prompt_token_ids": [10], + "token_ids": [11], + "logprobs": { + "tokens": ["token_id:11"], + "token_logprobs": [-0.1], + "top_logprobs": [{}], + "text_offset": [0], + }, + }, + { + **payload["choices"][0], + "index": 1, + "text": "two", + "prompt_token_ids": [20], + "token_ids": [21], + "logprobs": { + "tokens": ["token_id:21"], + "token_logprobs": [-0.2], + "top_logprobs": [{}], + "text_offset": [0], + }, + }, + ] + exchange.response = Completion.model_validate(payload) + + tokenized = art.Trajectory( + exchanges=TrajectoryExchanges(completions=[exchange]) + ).tokenize(multi_history=True) + + assert [history.token_ids for history in tokenized.histories] == [ + [10, 11], + [20, 21], + ] + assert [history.flags for history in tokenized.histories] == [ + [art.TokenFlag.EXACT, art.TokenFlag.EXACT | art.TokenFlag.SAMPLED], + [art.TokenFlag.EXACT, art.TokenFlag.EXACT | art.TokenFlag.SAMPLED], + ] + + +def test_history_tokenization_rejects_negative_source_indices() -> None: + exchange = _response_exchange("response-negative-source", 2) + history = art.Trajectory( + exchanges=TrajectoryExchanges(responses=[exchange]) + ).responses_history() + history.input_sources[-1] = ResponsesItemSource( + exchange=exchange, + output_index=-1, + generation_index=0, + ) + + with pytest.raises(ValueError, match="out of bounds"): + history.tokenize() + + +def test_batched_completions_reject_ambiguous_choice_indices() -> None: + exchange = _completion_exchange(prompt=["first", "second"]) + payload = exchange.response.model_dump(mode="python") + payload["choices"] = [ + {**payload["choices"][0], "index": 0}, + {**payload["choices"][0], "index": 2}, + ] + exchange.response = Completion.model_validate(payload) + + with pytest.raises(ValueError, match="Ambiguous"): + art.Trajectory(exchanges=TrajectoryExchanges(completions=[exchange])).tokenize( + multi_history=True + ) + + +def test_randomized_completions_projection_preserves_every_choice_once() -> None: + from art.trajectories._tokenize import _tokenize_trajectory_with_trace + + rng = random.Random(0) + for case in range(20): + prompt_count = rng.randint(1, 5) + choices_per_prompt = rng.randint(1, 4) + exchange = _completion_exchange( + prompt=[f"prompt-{case}-{index}" for index in range(prompt_count)] + ) + exchange.request["n"] = choices_per_prompt + template = exchange.response.model_dump(mode="python")["choices"][0] + choices: list[dict[str, object]] = [] + expected: list[list[int]] = [] + for prompt_index in range(prompt_count): + prompt_id = 10_000 + case * 100 + prompt_index + for local_choice in range(choices_per_prompt): + choice_index = prompt_index * choices_per_prompt + local_choice + output_id = 20_000 + case * 100 + choice_index + expected.append([prompt_id, output_id]) + choices.append( + { + **template, + "index": choice_index, + "text": f"answer-{output_id}", + "prompt_token_ids": [prompt_id], + "token_ids": [output_id], + "logprobs": { + "tokens": [f"token_id:{output_id}"], + "token_logprobs": [-0.1], + "top_logprobs": [{}], + "text_offset": [0], + }, + } + ) + rng.shuffle(choices) + response = exchange.response.model_dump(mode="python") + response["choices"] = choices + exchange.response = Completion.model_validate(response) + + trajectory = art.Trajectory( + exchanges=TrajectoryExchanges(completions=[exchange]) + ) + tokenized = trajectory.tokenize(multi_history=True) + + assert [history.token_ids for history in tokenized.histories] == expected + assert all( + history.flags + == [art.TokenFlag.EXACT, art.TokenFlag.EXACT | art.TokenFlag.SAMPLED] + for history in tokenized.histories + ) + traced, traces = _tokenize_trajectory_with_trace(trajectory) + assert [history.token_ids for history in traced.histories] == expected + assert [ + next(key for key in trace.source_keys if key is not None).prompt_index + for trace in traces + ] == [ + prompt_index + for prompt_index in range(prompt_count) + for _ in range(choices_per_prompt) + ] + assert all(len(trace.sources) == 1 for trace in traces) + + +def test_branching_and_multiple_models_require_explicit_resolution() -> None: + alternate = _chat_exchange([9], [3], offset=1) + alternate.request["messages"] = [{"role": "user", "content": "alternate"}] + branching = art.Trajectory( + exchanges=TrajectoryExchanges( + chat_completions=[ + _chat_exchange([1], [2], offset=0), + alternate, + ] + ) + ) + with pytest.raises(ValueError, match="exactly one history"): + branching.tokenize() + assert len(branching.tokenize(multi_history=True).histories) == 2 + + mixed = art.Trajectory( + exchanges=TrajectoryExchanges( + chat_completions=[ + _chat_exchange([1], [2], model="one", offset=0), + _chat_exchange([3], [4], model="two", offset=1), + ] + ) + ) + with pytest.raises(ValueError, match="exactly one model"): + mixed.tokenize() + assert mixed.tokenize(model="two").token_ids == [3, 4] + assert [ + history.model for history in mixed.tokenize(multi_history=True).histories + ] == ["one", "two"] + + +def test_model_selection_prefers_exact_identity_over_glob_interpretation() -> None: + literal_model = "org/model[1]" + wildcard_match = _chat_exchange([3], [4], model="org/model1", offset=1) + exact_match = _chat_exchange([1], [2], model=literal_model) + trajectory = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[wildcard_match, exact_match]) + ) + + tokenized = trajectory.tokenize(model=literal_model) + + assert tokenized.model == literal_model + assert tokenized.token_ids == [1, 2] + + +def test_legacy_additional_histories_require_multi_history_and_model() -> None: + first = _chat_exchange([1], [2]).response.choices[0] + second = _chat_exchange([3], [4]).response.choices[0] + trajectory = art.Trajectory( + messages_and_choices=[first], + additional_histories=[art.LegacyHistory(messages_and_choices=[second])], + ) + + with pytest.raises(ValueError, match="exactly one history"): + trajectory.tokenize(model="test/model") + with pytest.raises(ValueError, match="requires model="): + trajectory.tokenize(multi_history=True) + + tokenized = trajectory.tokenize(multi_history=True, model="test/model") + + assert [history.token_ids for history in tokenized.histories] == [[1, 2], [3, 4]] + + +class _FakeTokenizer: + def __init__(self) -> None: + self.calls: list[dict[str, object]] = [] + + def __call__(self, text: str, *, add_special_tokens: bool = False) -> list[int]: + del text, add_special_tokens + return [11] + + def apply_chat_template( + self, messages: list[dict[str, Any]], **kwargs: object + ) -> list[int]: + self.calls.append(kwargs) + return [10, 11] if messages[-1]["role"] == "assistant" else [10] + + +def test_fallback_uses_template_overrides_and_nan_logprobs( + monkeypatch: pytest.MonkeyPatch, +) -> None: + exchange = _message_exchange( + MessagesRequest( + model="wandb-artifact:///entity/project/run:step0", + messages=[{"role": "user", "content": "question"}], + chat_template="request-template", + chat_template_kwargs={"request": True}, + thinking={"type": "enabled", "budget_tokens": 128}, + ), + duration=timedelta(seconds=1), + ) + tokenizer = _FakeTokenizer() + loaded_base_models: list[str] = [] + monkeypatch.setattr( + "art.trajectories._tokenize._load_tokenizer", + lambda config: loaded_base_models.append(config.base_model) or tokenizer, + ) + monkeypatch.setattr( + "art.trajectories._tokenize._artifact_config", + lambda _model: pytest.fail("explicit base_model should bypass W&B"), + ) + + result = art.Trajectory( + exchanges=TrajectoryExchanges(messages=[exchange]) + ).tokenize( + base_model="base/model", + chat_template="explicit-template", + chat_template_kwargs={"explicit": True}, + ) + + assert result.token_ids == [10, 11] + assert loaded_base_models == ["base/model"] + assert result.flags == [art.TokenFlag(0), art.TokenFlag.SAMPLED] + assert math.isnan(result.logprobs[1]) + assert len(tokenizer.calls) == 3 + assert [call["add_generation_prompt"] for call in tokenizer.calls] == [ + False, + False, + True, + ] + assert [call["tokenize"] for call in tokenizer.calls] == [True, False, True] + assert all( + { + key: value + for key, value in call.items() + if key not in {"add_generation_prompt", "tokenize"} + } + == { + "tools": None, + "chat_template": "explicit-template", + "request": True, + "explicit": True, + "enable_thinking": True, + "thinking_budget": 128, + } + for call in tokenizer.calls + ) + + +@pytest.mark.parametrize( + ("model", "artifact_name"), + [ + ("wandb-artifact:///entity/project/run", "entity/project/run:latest"), + ("wandb-artifact:///entity/project/run:step0", "entity/project/run:step0"), + ], +) +def test_checkpoint_fallback_preserves_artifact_version_and_renderer( + monkeypatch: pytest.MonkeyPatch, + model: str, + artifact_name: str, +) -> None: + artifact_names: list[str] = [] + + class Api: + def artifact(self, name: str) -> SimpleNamespace: + artifact_names.append(name) + return SimpleNamespace( + metadata={ + "wandb.base_model": "base/model", + "renderer": { + "tokenizer_revision": "revision", + "chat_template": "template", + "chat_template_kwargs": {"thinking": True}, + }, + } + ) + + wandb = ModuleType("wandb") + apis = ModuleType("wandb.apis") + public = ModuleType("wandb.apis.public") + setattr(public, "Api", Api) + setattr(apis, "public", public) + setattr(wandb, "apis", apis) + monkeypatch.setitem(sys.modules, "wandb", wandb) + monkeypatch.setitem(sys.modules, "wandb.apis", apis) + monkeypatch.setitem(sys.modules, "wandb.apis.public", public) + exchange = _chat_exchange([], [], model=model) + extra = exchange.response.choices[0].model_extra + assert extra is not None + extra.pop("prompt_token_ids") + extra.pop("token_ids") + exchange.response.choices[0].logprobs = None + tokenizer = _FakeTokenizer() + configs = [] + monkeypatch.setattr( + "art.trajectories._tokenize._load_tokenizer", + lambda config: configs.append(config) or tokenizer, + ) + + art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[exchange]) + ).tokenize() + config = configs[0] + + assert artifact_names == [artifact_name] + assert config.base_model == "base/model" + assert config.revision == "revision" + assert config.chat_template == "template" + assert config.chat_template_kwargs == {"thinking": True} + assert tokenizer.calls[0]["chat_template"] == "template" + assert tokenizer.calls[0]["thinking"] is True + + +def test_loaded_tokenizers_are_cached_by_model_and_revision( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from art.trajectories._tokenize import ( + _cached_tokenizer, + _load_tokenizer, + _TokenizerConfig, + ) + + loaded: list[tuple[str, str | None]] = [] + + class AutoTokenizer: + @staticmethod + def from_pretrained(model: str, *, revision: str | None) -> object: + loaded.append((model, revision)) + return object() + + transformers = ModuleType("transformers") + setattr(transformers, "AutoTokenizer", AutoTokenizer) + monkeypatch.setitem(sys.modules, "transformers", transformers) + _cached_tokenizer.cache_clear() + try: + config = _TokenizerConfig("test/model", revision="revision") + assert _load_tokenizer(config) is _load_tokenizer(config) + assert loaded == [("test/model", "revision")] + finally: + _cached_tokenizer.cache_clear() + + +def test_deepseek_v4_uses_arts_protocol_renderer( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from art.trajectories._tokenize import _cached_tokenizer + + raw = object() + wrapped = object() + + class AutoTokenizer: + @staticmethod + def from_pretrained(model: str, *, revision: str | None) -> object: + assert model == "deepseek-ai/DeepSeek-V4-Flash" + assert revision is None + return raw + + transformers = ModuleType("transformers") + setattr(transformers, "AutoTokenizer", AutoTokenizer) + monkeypatch.setitem(sys.modules, "transformers", transformers) + monkeypatch.setattr( + "art.megatron.dsv4.tokenizer.get_dsv4_tokenizer", + lambda tokenizer: ( + wrapped if tokenizer is raw else pytest.fail("wrong tokenizer") + ), + ) + _cached_tokenizer.cache_clear() + try: + assert _cached_tokenizer("deepseek-ai/DeepSeek-V4-Flash", None) is wrapped + finally: + _cached_tokenizer.cache_clear() + + +def test_anthropic_fallback_preserves_thinking_and_tool_history() -> None: from art.trajectories._tokenize import _anthropic_messages - messages = _anthropic_messages( + messages = _anthropic_messages( + { + "system": [{"type": "text", "text": "system"}], + "messages": [ + {"role": "user", "content": "question"}, + { + "role": "assistant", + "content": [ + {"type": "thinking", "thinking": "reason"}, + {"type": "text", "text": "calling"}, + { + "type": "tool_use", + "id": "call-1", + "name": "lookup", + "input": {"key": "value"}, + }, + ], + }, + { + "role": "user", + "content": [ + { + "type": "tool_result", + "tool_use_id": "call-1", + "content": [{"type": "text", "text": "result"}], + }, + {"type": "text", "text": "continue"}, + ], + }, + ], + } + ) + + assert messages == [ + {"role": "system", "content": "system"}, + {"role": "user", "content": "question"}, + { + "role": "assistant", + "content": "calling", + "reasoning": "reason", + "tool_calls": [ + { + "id": "call-1", + "type": "function", + "function": { + "name": "lookup", + "arguments": '{"key": "value"}', + }, + } + ], + }, + {"role": "tool", "tool_call_id": "call-1", "content": "result"}, + {"role": "user", "content": "continue"}, + ] + + +@pytest.mark.parametrize("top_level_only", [False, True]) +def test_reasoning_stripped_messages_history_preserves_exact_tokens( + top_level_only: bool, +) -> None: + def exchange( + offset: int, + request_messages: list[MessageParam], + answer: str, + prompt_token_ids: list[int], + token_ids: list[int], + ) -> MessagesExchange: + start = datetime(2026, 1, 1) + timedelta(seconds=offset) + thinking: dict[str, Any] = { + "type": "thinking", + "thinking": f"thought-{offset}", + "signature": "signature", + } + text: dict[str, Any] = {"type": "text", "text": answer} + response_extra: dict[str, Any] = {} + if top_level_only: + response_extra = { + "prompt_token_ids": prompt_token_ids, + "token_ids": [90 + offset, *token_ids], + "logprobs": [ + -9.0 - offset, + *[-0.1 * token for token in token_ids], + ], + } + else: + thinking.update({"token_ids": [90 + offset], "logprobs": [-9.0 - offset]}) + text.update( + { + "token_ids": token_ids, + "logprobs": [-0.1 * token for token in token_ids], + } + ) + return MessagesExchange( + request=MessagesRequest( + model="test/model", + messages=request_messages, + max_tokens=16, + thinking={"type": "enabled", "budget_tokens": 8}, + ), + response=Message.model_validate( + { + "id": f"message-{offset}", + "type": "message", + "role": "assistant", + "model": "test/model", + "content": [thinking, text], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 1, "output_tokens": len(token_ids)}, + **response_extra, + } + ), + start_time=start, + end_time=start + timedelta(milliseconds=1), + ) + + first = exchange( + 0, + [{"role": "user", "content": "one"}], + "first", + [10], + [101, 102], + ) + second = exchange( + 1, + [ + {"role": "user", "content": "one"}, + {"role": "assistant", "content": "first"}, + {"role": "user", "content": "two"}, + ], + "second", + [10, 101, 102, 11], + [201], + ) + + class Tokenizer: + def __call__(self, text: str, *, add_special_tokens: bool = False) -> list[int]: + assert not add_special_tokens + return { + "one": [10], + "first": [50], + "second": [60], + "thought-1": [70], + "two": [11], + }[text] + + def apply_chat_template( + self, + messages: list[dict[str, Any]], + **kwargs: object, + ) -> list[int]: + del kwargs + by_length = { + 1: [10], + 2: [10, 50], + 3: [10, 50, 11], + 4: [10, 50, 11, 70, 60], + } + return by_length[len(messages)] + + history = art.Trajectory( + exchanges=TrajectoryExchanges(messages=[first, second]) + ).anthropic_messages_histories()[1] + tokenized = history.tokenize(tokenizer=Tokenizer()) + + assert tokenized.token_ids == [10, 101, 102, 11, 91, 201] + assert tokenized.logprobs[1:3] == pytest.approx([-10.1, -10.2]) + assert tokenized.logprobs[-2] == pytest.approx(-10.0) + assert tokenized.logprobs[-1] == pytest.approx(-20.1) + assert tokenized.flags[1:3] == [ + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + ] + assert tokenized.flags[-2:] == [ + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + ] + + +def test_choice_logprobs_survive_tokenizer_fallback( + monkeypatch: pytest.MonkeyPatch, +) -> None: + exchange = _chat_exchange([], []) + logprobs = exchange.response.choices[0].logprobs + assert logprobs is not None + exchange.response.choices[0].logprobs = logprobs.model_copy( + update={ + "content": [ + ChatCompletionTokenLogprob( + token="answer", + logprob=-0.7, + bytes=list(b"answer"), + top_logprobs=[], + ) + ] + } + ) + + class Tokenizer: + def apply_chat_template( + self, messages: list[dict[str, Any]], **kwargs: object + ) -> list[int]: + del kwargs + if messages[-1]["role"] != "assistant": + return [10] + if str(messages[-1]["content"]).startswith("ART_TRAJECTORY_"): + return [10, 99, 12] + return [10, 11, 12] + + def __call__(self, text: str, **kwargs: object) -> SimpleNamespace: + del kwargs + return SimpleNamespace( + input_ids={"turn 0": [10], "answer": [11], "turn 1": [20]}[text] + ) + + monkeypatch.setattr( + "art.trajectories._tokenize._load_tokenizer", lambda _config: Tokenizer() + ) + result = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[exchange]) + ).tokenize( + base_model="base/model", + ) + assert result.token_ids == [10, 11, 12] + assert result.logprobs[1] == -0.7 + assert math.isnan(result.logprobs[2]) + + +def test_chat_fallback_rejects_unique_text_match_in_generation_scaffold() -> None: + exchange = _chat_exchange([], []) + exchange.request["messages"] = [{"role": "user", "content": "question"}] + choice = exchange.response.choices[0] + assert choice.model_extra is not None + choice.model_extra.pop("prompt_token_ids", None) + choice.model_extra.pop("token_ids", None) + assert choice.logprobs is not None + choice.logprobs = choice.logprobs.model_copy( + update={ + "content": [ + ChatCompletionTokenLogprob( + token="answer", + logprob=-0.7, + bytes=list(b"answer"), + top_logprobs=[], + ) + ] + } + ) + + class Tokenizer: + def apply_chat_template( + self, + messages: list[dict[str, Any]], + *, + add_generation_prompt: bool, + **kwargs: object, + ) -> list[int]: + del kwargs + if messages[-1]["role"] == "assistant": + return [1, 11, 12] + return [1, 11] if add_generation_prompt else [1] + + def __call__(self, text: str, **kwargs: object) -> SimpleNamespace: + del kwargs + return SimpleNamespace(input_ids={"question": [1], "answer": [11]}[text]) + + history = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[exchange]) + ).chat_completions_history() + + with pytest.raises(ValueError, match="uniquely locate"): + history.tokenize(tokenizer=Tokenizer()) + + +def test_chat_exact_ids_reject_unique_match_in_generation_scaffold() -> None: + exchange = _chat_exchange([], [11]) + exchange.request["messages"] = [{"role": "user", "content": "question"}] + + class Tokenizer: + def apply_chat_template( + self, + messages: list[dict[str, Any]], + *, + add_generation_prompt: bool, + **kwargs: object, + ) -> list[int]: + del kwargs + if messages[-1]["role"] == "assistant": + return [1, 11, 12] + return [1, 11] if add_generation_prompt else [1] + + def __call__(self, text: str, **kwargs: object) -> SimpleNamespace: + del kwargs + return SimpleNamespace(input_ids={"question": [1], "answer": [11]}[text]) + + history = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[exchange]) + ).chat_completions_history() + + with pytest.raises(ValueError, match="uniquely locate"): + history.tokenize(tokenizer=Tokenizer()) + + +def test_chat_prompt_ids_do_not_bind_output_to_trailing_scaffold() -> None: + exchange = _chat_exchange([1, 2], []) + choice = exchange.response.choices[0] + assert choice.logprobs is not None + choice.logprobs = choice.logprobs.model_copy( + update={ + "content": [ + ChatCompletionTokenLogprob( + token="answer", + logprob=-0.7, + bytes=list(b"answer"), + top_logprobs=[], + ) + ] + } + ) + + class Tokenizer: + def apply_chat_template( + self, messages: list[dict[str, Any]], **kwargs: object + ) -> list[int]: + del kwargs + return [1, 2, 12, 11] if messages[-1]["role"] == "assistant" else [1, 2] + + def __call__(self, text: str, **kwargs: object) -> SimpleNamespace: + del kwargs + return SimpleNamespace( + input_ids={"turn 0": [10], "answer": [11], "turn 1": [20]}[text] + ) + + history = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[exchange]) + ).chat_completions_history() + + with pytest.raises(ValueError, match="content boundary"): + history.tokenize(tokenizer=Tokenizer()) + + +def test_chat_exact_hidden_suffix_preserves_rendered_trailing_scaffold() -> None: + exchange = _chat_exchange([1, 99], [7, 8]) + history = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[exchange]) + ).chat_completions_history() + history.chat_template = "rerender" + + class Tokenizer: + def apply_chat_template( + self, messages: list[dict[str, Any]], **kwargs: object + ) -> list[int]: + del kwargs + if messages[-1]["role"] != "assistant": + return [1, 99] + if str(messages[-1]["content"]).startswith("ART_TRAJECTORY_"): + return [1, 99, 999, 9] + return [1, 99, 7, 9] + + def __call__(self, text: str, **kwargs: object) -> SimpleNamespace: + del text, kwargs + return SimpleNamespace(input_ids=[7]) + + tokenized = history.tokenize(tokenizer=Tokenizer()) + + assert tokenized.token_ids == [1, 99, 7, 8, 9] + assert tokenized.flags == [ + art.TokenFlag(0), + art.TokenFlag(0), + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + art.TokenFlag(0), + ] + + +def test_each_chat_choice_preserves_its_visible_fallback_logprobs() -> None: + exchange = _chat_exchange([], []) + exchange.request["messages"] = [{"role": "user", "content": "question"}] + data = exchange.response.model_dump(mode="python") + choices = [] + for index, (text, logprob) in enumerate((("left", -0.1), ("right", -0.2))): + choice = data["choices"][0].copy() + choice.pop("prompt_token_ids", None) + choice.pop("token_ids", None) + choice["index"] = index + choice["message"] = {"role": "assistant", "content": text} + choice["logprobs"] = { + "content": [ + { + "token": text, + "logprob": logprob, + "bytes": list(text.encode()), + "top_logprobs": [], + } + ] + } + choices.append(choice) + data["choices"] = choices + exchange.response = ChatCompletion.model_validate(data) + trajectory = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[exchange]) + ) + + class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> list[int]: + del kwargs + return {"question": [1], "left": [2], "right": [3]}[text] + + def apply_chat_template( + self, messages: list[dict[str, Any]], **kwargs: object + ) -> list[int]: + del kwargs + return [ + token for message in messages for token in self(str(message["content"])) + ] + + histories = trajectory.chat_completions_histories() + direct = [history.tokenize(tokenizer=Tokenizer()) for history in histories] + tokenized = trajectory.tokenize(multi_history=True, tokenizer=Tokenizer()) + + assert [history.logprobs[-1] for history in direct] == [-0.1, -0.2] + assert [history.logprobs[-1] for history in tokenized.histories] == [-0.1, -0.2] + + from art.trajectories._tokenize import _tokenize_trajectory_with_trace + + traced, traces = _tokenize_trajectory_with_trace(trajectory, tokenizer=Tokenizer()) + assert [history.logprobs[-1] for history in traced.histories] == [-0.1, -0.2] + assert [ + {key.index for key in trace.source_keys if key is not None} for trace in traces + ] == [{0}, {1}] + + +def test_chat_fallback_anchors_sampled_text_away_from_equal_user_text() -> None: + first = _chat_exchange([], [], offset=0) + first.request["messages"] = [{"role": "user", "content": "q"}] + first_data = first.response.model_dump(mode="python") + first_choice = first_data["choices"][0] + first_choice.pop("prompt_token_ids", None) + first_choice.pop("token_ids", None) + first_choice["message"]["content"] = "same" + first_choice["logprobs"]["content"] = [ + { + "token": "same", + "logprob": -0.1, + "bytes": list(b"same"), + "top_logprobs": [], + } + ] + first.response = ChatCompletion.model_validate(first_data) + + second = _chat_exchange([], [], offset=1) + second.request["messages"] = [ + {"role": "user", "content": "q"}, + {"role": "assistant", "content": "same"}, + {"role": "user", "content": "same"}, + ] + second_data = second.response.model_dump(mode="python") + second_choice = second_data["choices"][0] + second_choice.pop("prompt_token_ids", None) + second_choice.pop("token_ids", None) + second_choice["message"]["content"] = "other" + second_choice["logprobs"]["content"] = [ + { + "token": "other", + "logprob": -0.2, + "bytes": list(b"other"), + "top_logprobs": [], + } + ] + second.response = ChatCompletion.model_validate(second_data) + + class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> list[int]: + del kwargs + return {"q": [1], "same": [2], "other": [3]}[text] + + def apply_chat_template( + self, messages: list[dict[str, Any]], **kwargs: object + ) -> list[int]: + del kwargs + return [ + token for message in messages for token in self(str(message["content"])) + ] + + history = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[first, second]) + ).chat_completions_history() + tokenized = history.tokenize(tokenizer=Tokenizer()) + + assert tokenized.token_ids == [1, 2, 2, 3] + assert tokenized.flags == [ + art.TokenFlag(0), + art.TokenFlag.SAMPLED, + art.TokenFlag(0), + art.TokenFlag.SAMPLED, + ] + assert tokenized.logprobs[1] == -0.1 + assert math.isnan(tokenized.logprobs[2]) + assert tokenized.logprobs[3] == -0.2 + + +def test_exact_chat_ids_accept_ordinary_positional_logprobs() -> None: + exchange = _chat_exchange([1], [2]) + logprobs = exchange.response.choices[0].logprobs + assert logprobs is not None and logprobs.content + exchange.response.choices[0].logprobs = logprobs.model_copy( + update={ + "content": [ + logprobs.content[0].model_copy( + update={"token": "answer", "logprob": -0.75} + ) + ] + } + ) + + tokenized = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[exchange]) + ).tokenize() + + assert tokenized.token_ids == [1, 2] + assert tokenized.logprobs[1] == -0.75 + + +def test_chat_view_preserves_two_turn_textual_logprobs_without_exact_ids() -> None: + exchanges = [_chat_exchange([], [], offset=index) for index in range(2)] + for index, exchange in enumerate(exchanges): + choice = exchange.response.choices[0] + assert choice.model_extra is not None + choice.model_extra.pop("prompt_token_ids", None) + choice.model_extra.pop("token_ids", None) + assert choice.logprobs is not None + choice.logprobs = choice.logprobs.model_copy( + update={ + "content": [ + ChatCompletionTokenLogprob( + token="answer", + logprob=-0.4 - index / 10, + bytes=list(b"answer"), + top_logprobs=[], + ) + ] + } + ) + + class Tokenizer: + def apply_chat_template( + self, messages: list[dict[str, Any]], **kwargs: object + ) -> list[int]: + del kwargs + if len(messages) == 4: + return [10, 11, 20, 11] + if messages[-1]["role"] == "assistant": + return [10, 11] + return [10] + + def __call__(self, text: str, **kwargs: object) -> SimpleNamespace: + del kwargs + return SimpleNamespace( + input_ids={"turn 0": [10], "answer": [11], "turn 1": [20]}[text] + ) + + history = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=exchanges) + ).chat_completions_history() + tokenized = history.tokenize(tokenizer=Tokenizer()) + + assert tokenized.token_ids == [10, 11, 20, 11] + assert tokenized.logprobs[1] == pytest.approx(-0.4) + assert tokenized.logprobs[3] == pytest.approx(-0.5) + assert tokenized.flags == [ + art.TokenFlag(0), + art.TokenFlag.SAMPLED, + art.TokenFlag(0), + art.TokenFlag.SAMPLED, + ] + + +def test_chat_view_uses_later_exact_prompt_when_first_is_missing() -> None: + first = _chat_exchange([], [11]) + second = _chat_exchange([10, 11, 20], [], offset=1) + second.response.choices[0].message.content = "final" + choice = second.response.choices[0] + assert choice.model_extra is not None + choice.model_extra.pop("token_ids", None) + assert choice.logprobs is not None + choice.logprobs = choice.logprobs.model_copy( + update={ + "content": [ + ChatCompletionTokenLogprob( + token="final", + logprob=-0.5, + bytes=list(b"final"), + top_logprobs=[], + ) + ] + } + ) + + class Tokenizer: + def apply_chat_template( + self, messages: list[dict[str, Any]], **kwargs: object + ) -> list[int]: + del kwargs + token_ids = { + "turn 0": [10], + "answer": [11], + "turn 1": [20], + "final": [21], + } + return [ + token + for message in messages + for token in token_ids[str(message["content"])] + ] + + def __call__(self, text: str, **kwargs: object) -> SimpleNamespace: + del kwargs + return SimpleNamespace( + input_ids={ + "turn 0": [10], + "answer": [11], + "turn 1": [20], + "final": [21], + }[text] + ) + + history = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[first, second]) + ).chat_completions_history() + tokenized = history.tokenize(tokenizer=Tokenizer()) + + assert tokenized.token_ids == [10, 11, 20, 21] + assert tokenized.logprobs[-1] == pytest.approx(-0.5) + assert tokenized.flags == [ + art.TokenFlag.EXACT, + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + art.TokenFlag.EXACT, + art.TokenFlag.SAMPLED, + ] + + +def test_chat_content_and_refusal_logprobs_are_combined_in_protocol_order() -> None: + exchange = _chat_exchange([1], [2, 3]) + data = exchange.response.model_dump(mode="python") + choice = data["choices"][0] + choice["message"]["refusal"] = "refusal" + choice["logprobs"] = { + "content": [ + { + "token": "token_id:2", + "logprob": -0.2, + "bytes": [], + "top_logprobs": [], + } + ], + "refusal": [ + { + "token": "token_id:3", + "logprob": -0.3, + "bytes": [], + "top_logprobs": [], + } + ], + } + exchange.response = ChatCompletion.model_validate(data) + + tokenized = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[exchange]) + ).tokenize() + + assert tokenized.token_ids == [1, 2, 3] + assert tokenized.logprobs[1:] == pytest.approx([-0.2, -0.3]) + + +def test_chat_content_and_refusal_logprobs_must_match_exact_ids() -> None: + exchange = _chat_exchange([1], [2, 3]) + data = exchange.response.model_dump(mode="python") + choice = data["choices"][0] + choice["message"]["refusal"] = "refusal" + choice["logprobs"] = { + "content": [ + { + "token": "token_id:2", + "logprob": -0.2, + "bytes": [], + "top_logprobs": [], + } + ], + "refusal": [ + { + "token": "token_id:4", + "logprob": -0.4, + "bytes": [], + "top_logprobs": [], + } + ], + } + exchange.response = ChatCompletion.model_validate(data) + + with pytest.raises(ValueError, match="disagree with choice logprobs"): + art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[exchange]) + ).tokenize() + + +def test_chat_visible_logprobs_include_content_and_refusal() -> None: + from art.trajectories._tokenize import _visible_logprobs + + exchange = _chat_exchange([], []) + data = exchange.response.model_dump(mode="python") + choice = data["choices"][0] + choice["message"]["refusal"] = "refusal" + choice["logprobs"] = { + "content": [ + { + "token": "answer", + "logprob": -0.2, + "bytes": list(b"answer"), + "top_logprobs": [], + } + ], + "refusal": [ + { + "token": "refusal", + "logprob": -0.3, + "bytes": list(b"refusal"), + "top_logprobs": [], + } + ], + } + exchange.response = ChatCompletion.model_validate(data) + + assert _visible_logprobs(exchange) == [ + ("answer", -0.2), + ("refusal", -0.3), + ] + + +def test_empty_chat_prompt_ids_are_missing_evidence( + monkeypatch: pytest.MonkeyPatch, +) -> None: + exchange = _chat_exchange([], [2]) + + class Tokenizer: + def apply_chat_template( + self, messages: list[dict[str, Any]], **kwargs: object + ) -> list[int]: + del kwargs + return [10, 20] if messages[-1]["role"] == "assistant" else [10] + + def __call__(self, text: str, **kwargs: object) -> list[int]: + del text, kwargs + return [20] + + monkeypatch.setattr( + "art.trajectories._tokenize._load_tokenizer", lambda _config: Tokenizer() + ) + tokenized = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[exchange]) + ).tokenize(base_model="base/model") + + assert tokenized.token_ids == [10, 2] + assert tokenized.flags == [ + art.TokenFlag(0), + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + ] + + +def test_missing_completion_renders_only_missing_region_when_prompt_is_exact( + monkeypatch: pytest.MonkeyPatch, +) -> None: + exchange = _chat_exchange([99], []) + extra = exchange.response.choices[0].model_extra + assert extra is not None + extra.pop("token_ids", None) + logprobs = exchange.response.choices[0].logprobs + assert logprobs is not None + exchange.response.choices[0].logprobs = logprobs.model_copy( + update={ + "content": [ + ChatCompletionTokenLogprob( + token="answer", + logprob=-0.5, + bytes=list(b"answer"), + top_logprobs=[], + ) + ] + } + ) + + class Tokenizer: + def apply_chat_template( + self, messages: list[dict[str, Any]], **kwargs: object + ) -> list[int]: + del kwargs + return [10, 20] if messages[-1]["role"] == "assistant" else [10] + + def __call__(self, text: str, **kwargs: object) -> list[int]: + del text, kwargs + return [20] + + monkeypatch.setattr( + "art.trajectories._tokenize._load_tokenizer", lambda _config: Tokenizer() + ) + tokenized = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[exchange]) + ).tokenize(base_model="base/model") + + assert tokenized.token_ids == [99, 20] + assert tokenized.logprobs[1] == -0.5 + assert tokenized.flags == [ + art.TokenFlag.EXACT, + art.TokenFlag.SAMPLED, + ] + + +def test_ambiguous_visible_logprobs_raise( + monkeypatch: pytest.MonkeyPatch, +) -> None: + exchange = _chat_exchange([], []) + logprobs = exchange.response.choices[0].logprobs + assert logprobs is not None + exchange.response.choices[0].logprobs = logprobs.model_copy( + update={ + "content": [ + ChatCompletionTokenLogprob( + token="answer", + logprob=-0.7, + bytes=list(b"answer"), + top_logprobs=[], + ) + ] + } + ) + + class Tokenizer: + def apply_chat_template( + self, messages: list[dict[str, Any]], **kwargs: object + ) -> list[int]: + del kwargs + return [10, 11, 12, 11] if messages[-1]["role"] == "assistant" else [10] + + def __call__(self, text: str, **kwargs: object) -> SimpleNamespace: + del text, kwargs + return SimpleNamespace(input_ids=[11]) + + monkeypatch.setattr( + "art.trajectories._tokenize._load_tokenizer", lambda _config: Tokenizer() + ) + with pytest.raises(ValueError, match="uniquely locate"): + art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[exchange]) + ).tokenize( + base_model="base/model", + ) + + +def test_legacy_token_and_logprob_length_mismatch_raises() -> None: + exchange = _chat_exchange([1], [2, 3]) + choice = exchange.response.choices[0] + assert choice.logprobs is not None + content = choice.logprobs.content + assert content + choice.logprobs = choice.logprobs.model_copy( + update={ + "content": [ + content[0].model_copy( + update={"token": "answer", "bytes": list(b"answer")} + ) + ] + } + ) + + with pytest.raises(ValueError, match="differ in length"): + art.Trajectory(messages_and_choices=[choice]).tokenize(model="test/model") + + +def test_anthropic_fallback_rejects_unknown_content_blocks( + monkeypatch: pytest.MonkeyPatch, +) -> None: + image: ImageBlockParam = { + "type": "image", + "source": { + "type": "base64", + "media_type": "image/png", + "data": "...", + }, + } + message: MessageParam = {"role": "user", "content": [image]} + exchange = _message_exchange( + MessagesRequest( + model="test/model", + messages=[message], + ), + duration=timedelta(seconds=1), + ) + monkeypatch.setattr( + "art.trajectories._tokenize._load_tokenizer", lambda _config: _FakeTokenizer() + ) + + with pytest.raises(ValueError, match="Unsupported Anthropic content block"): + art.Trajectory(exchanges=TrajectoryExchanges(messages=[exchange])).tokenize( + base_model="base/model", + ) + + +def test_undecodable_visible_token_bytes_fall_back_to_nan( + monkeypatch: pytest.MonkeyPatch, +) -> None: + exchange = _chat_exchange([], []) + logprobs = exchange.response.choices[0].logprobs + assert logprobs is not None + exchange.response.choices[0].logprobs = logprobs.model_copy( + update={ + "content": [ + ChatCompletionTokenLogprob( + token="ordinary-token", + logprob=-0.7, + bytes=[0xF0], + top_logprobs=[], + ) + ] + } + ) + + class Tokenizer: + def apply_chat_template( + self, messages: list[dict[str, Any]], **kwargs: object + ) -> list[int]: + del kwargs + return [10, 11] if messages[-1]["role"] == "assistant" else [10] + + def __call__(self, text: str, **kwargs: object) -> list[int]: + del text, kwargs + return [11] + + monkeypatch.setattr( + "art.trajectories._tokenize._load_tokenizer", lambda _config: Tokenizer() + ) + result = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[exchange]) + ).tokenize( + base_model="base/model", + ) + + assert result.token_ids == [10, 11] + assert math.isnan(result.logprobs[1]) + + +def test_json_round_trip_preserves_exchange_types() -> None: + exchange = _chat_exchange([1], [2]) + request: dict[str, Any] = { + "model": "test/model", + "messages": [ + {"role": "assistant", "content": "answer", "reasoning": "thinking"} + ], + } + exchange.request = ChatCompletionsRequest(**request) + original = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[exchange]) + ) + dumped = original.model_dump(mode="json", warnings="error") + assert dumped["exchanges"]["chat_completions"][0]["request"] == request + restored = art.Trajectory.model_validate_json(original.model_dump_json()) + assert restored.model_dump(mode="json") == original.model_dump(mode="json") + assert isinstance(restored.exchanges.chat_completions[0].response, ChatCompletion) + + +def _response_exchange( + response_id: str, + output_id: int, + *, + previous_response_id: str | None = None, + offset: int = 0, + prompt_token_ids: list[int] | None = None, +) -> ResponsesExchange: + response = Response.model_validate( + { + "id": response_id, + "created_at": float(offset), + "model": "test/model", + "object": "response", + "output": [ + { + "id": f"message-{response_id}", + "type": "message", + "role": "assistant", + "status": "completed", + "content": [ + { + "type": "output_text", + "text": "answer", + "annotations": [], + "logprobs": [], + } + ], + } + ], + "parallel_tool_calls": True, + "tool_choice": "auto", + "tools": [], + "token_generations": [ + { + "prompt_token_ids": prompt_token_ids or [10], + "output_tokens": [{"token_id": output_id, "logprob": -0.1}], + "output_indices": [0], + } + ], + } + ) + request = ResponsesRequest(model="test/model", input=f"turn {offset}") + if previous_response_id is not None: + request["previous_response_id"] = previous_response_id + start = datetime(2026, 1, 1) + timedelta(seconds=offset) + return ResponsesExchange( + request=request, + response=response, + start_time=start, + end_time=start + timedelta(milliseconds=1), + ) + + +def test_cross_exchange_responses_reasoning_split_uses_later_prompt_backbone() -> None: + first = _response_exchange("response-1", 3, prompt_token_ids=[1]) + first_data = first.response.model_dump(mode="python") + first_data["output"] = [ + { + "id": "reasoning-response-1", + "type": "reasoning", + "summary": [{"type": "summary_text", "text": "think"}], + }, + first_data["output"][0], + ] + first_data["token_generations"] = [ + { + "prompt_token_ids": [1], + "output_tokens": [ + {"token_id": 2, "logprob": -0.2}, + {"token_id": 3, "logprob": -0.3}, + ], + "output_indices": [0, 1], + } + ] + first.response = Response.model_validate(first_data) + second = _response_exchange( + "response-2", + 5, + previous_response_id="response-1", + offset=1, + prompt_token_ids=[1, 3, 4], + ) + trajectory = art.Trajectory( + exchanges=TrajectoryExchanges(responses=[first, second]) + ) + + tokenized = trajectory.tokenize(multi_history=True) + + assert [history.token_ids for history in tokenized.histories] == [ + [1, 2, 3], + [1, 3, 4, 5], + ] + assert math.isnan(tokenized.histories[1].logprobs[0]) + assert tokenized.histories[1].logprobs[1] == -0.3 + assert math.isnan(tokenized.histories[1].logprobs[2]) + assert tokenized.histories[1].logprobs[3] == -0.1 + assert tokenized.histories[1].flags == [ + art.TokenFlag.EXACT, + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + art.TokenFlag.EXACT, + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + ] + + +def _response_with_content_logprobs(*, exact_second: bool) -> ResponsesExchange: + exchange = _response_exchange("response-content-logprobs", 0) + data = exchange.response.model_dump(mode="python") + data.pop("token_generations", None) + + def entry(token: str, token_id: int | None, logprob: float) -> dict[str, Any]: + return { + "token": token, + "logprob": logprob, + "bytes": list(("a" if token_id == 11 else "b").encode()), + "top_logprobs": [], + **({"token_id": token_id} if token_id is not None else {}), + } + + data["output"][0]["content"] = [ + { + "type": "output_text", + "text": "a", + "annotations": [], + "logprobs": [entry("token_id:11", 11, -0.1)], + }, + { + "type": "output_text", + "text": "b", + "annotations": [], + "logprobs": [ + entry( + "token_id:12" if exact_second else "b", + 12 if exact_second else None, + -0.2, + ) + ], + }, + ] + exchange.response = Response.model_validate(data) + return exchange + + +def test_responses_aggregates_complete_exact_pairs_across_content_blocks( + monkeypatch: pytest.MonkeyPatch, +) -> None: + class Tokenizer: + def apply_chat_template( + self, messages: list[dict[str, Any]], **kwargs: object + ) -> list[int]: + del messages, kwargs + return [10] + + monkeypatch.setattr( + "art.trajectories._tokenize._load_tokenizer", lambda _config: Tokenizer() + ) + result = art.Trajectory( + exchanges=TrajectoryExchanges( + responses=[_response_with_content_logprobs(exact_second=True)] + ) + ).tokenize( + base_model="base/model", + ) + + assert result.token_ids == [10, 11, 12] + assert result.logprobs[1:] == [-0.1, -0.2] + assert result.flags[1:] == [ + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + ] + + +def test_responses_tool_output_source_is_not_sampled_without_generation() -> None: + from art.trajectories._tokenize import ( + _responses_output_is_sampled, + _source_is_sampled, + ) + + exchange = _response_exchange("response-tool-output", 0) + data = exchange.response.model_dump(mode="python") + data.pop("token_generations", None) + data["output"] = [ + { + "type": "function_call_output", + "id": "output-1", + "call_id": "call-1", + "output": "result", + "status": "completed", + } + ] + exchange.response = Response.model_validate(data) + + history = ( + art.Trajectory(exchanges=TrajectoryExchanges(responses=[exchange])) + .responses_history() + .as_chat_completions_history() + ) + assert history.messages[-1]["role"] == "tool" + source = history.message_sources[-1] + assert source is not None + assert not _source_is_sampled(source) + assert not _responses_output_is_sampled({"type": "function_call_output"}) + + +def test_responses_source_rejects_boolean_generation_index() -> None: + from art.trajectories._tokenize import _responses_source_generation + + exchange = _response_exchange("response-bool-generation", 2) + source = ChatCompletionsMessageSource( + exchange=exchange, + output_indices=(0,), + generation_index=True, + ) + + with pytest.raises(ValueError, match="generation index is invalid"): + _responses_source_generation(source) + + +def test_responses_missing_token_generations_falls_back_for_visible_output( + monkeypatch: pytest.MonkeyPatch, +) -> None: + exchange = _response_exchange("response-empty-raw", 0) + data = exchange.response.model_dump(mode="python") + data.pop("token_generations", None) + exchange.response = Response.model_validate(data) + monkeypatch.setattr( + "art.trajectories._tokenize._load_tokenizer", lambda _config: _FakeTokenizer() + ) + + result = art.Trajectory( + exchanges=TrajectoryExchanges(responses=[exchange]) + ).tokenize( + base_model="base/model", + chat_template="template", + chat_template_kwargs={}, + ) + + assert result.token_ids == [10, 11] + + +def test_responses_does_not_use_partial_exact_content_pairs( + monkeypatch: pytest.MonkeyPatch, +) -> None: + class Tokenizer: + def apply_chat_template( + self, messages: list[dict[str, Any]], **kwargs: object + ) -> list[int]: + del kwargs + return [10, 11, 12] if messages[-1]["role"] == "assistant" else [10] + + def __call__(self, text: str, **kwargs: object) -> SimpleNamespace: + del kwargs + return SimpleNamespace( + input_ids=[11 if text in {"a", "token_id:11"} else 12] + ) + + monkeypatch.setattr( + "art.trajectories._tokenize._load_tokenizer", lambda _config: Tokenizer() + ) + result = art.Trajectory( + exchanges=TrajectoryExchanges( + responses=[_response_with_content_logprobs(exact_second=False)] + ) + ).tokenize( + base_model="base/model", + ) + + assert result.token_ids == [10, 11, 12] + assert result.logprobs[1:] == [-0.1, -0.2] + + +def test_responses_rejects_only_unrenderable_prompt_history( + monkeypatch: pytest.MonkeyPatch, +) -> None: + class Tokenizer: + def apply_chat_template( + self, messages: list[dict[str, Any]], **kwargs: object + ) -> list[int]: + del kwargs + assistant_count = sum( + message["role"] == "assistant" for message in messages + ) + return [10, *range(2, 2 + assistant_count)] + + monkeypatch.setattr( + "art.trajectories._tokenize._load_tokenizer", lambda _config: Tokenizer() + ) + request_reasoning = _response_exchange("request-reasoning", 2) + request_data = request_reasoning.response.model_dump(mode="python") + request_data.pop("token_generations", None) + request_reasoning.response = Response.model_validate(request_data) + request_reasoning.request["input"] = [ + { + "id": "reasoning-1", + "summary": [{"type": "summary_text", "text": "request thought"}], + "type": "reasoning", + } + ] + + response_reasoning = _response_exchange("response-reasoning", 2) + data = response_reasoning.response.model_dump(mode="python") + data["output"] = [ + { + "id": "reasoning-2", + "summary": [{"type": "summary_text", "text": "response thought"}], + "type": "reasoning", + } + ] + data.pop("token_generations", None) + response_reasoning.response = Response.model_validate(data) + + art.Trajectory( + exchanges=TrajectoryExchanges(responses=[request_reasoning]) + ).tokenize( + base_model="base/model", + ) + + single = art.Trajectory( + exchanges=TrajectoryExchanges(responses=[response_reasoning]) + ) + assert single.tokenize(base_model="base/model").token_ids == [ + 10, + 2, + ] + + continuation = _response_exchange( + "continuation", + 3, + previous_response_id=response_reasoning.response.id, + offset=1, + ) + continuation_data = continuation.response.model_dump(mode="python") + continuation_data.pop("token_generations", None) + continuation.response = Response.model_validate(continuation_data) + assert art.Trajectory( + exchanges=TrajectoryExchanges(responses=[response_reasoning, continuation]) + ).tokenize( + base_model="base/model", + ).token_ids == [10, 2, 3] + + +def test_responses_opaque_reasoning_requires_exact_tokens( + monkeypatch: pytest.MonkeyPatch, +) -> None: + exchange = _response_exchange("opaque-reasoning", 2) + response = exchange.response.model_dump(mode="python") + response["output"] = [ + { + "id": "reasoning-1", + "encrypted_content": "opaque", + "summary": [], + "type": "reasoning", + } + ] + exchange.response = Response.model_validate(response) + monkeypatch.setattr( + "art.trajectories._tokenize._load_tokenizer", lambda _config: _FakeTokenizer() + ) + + assert art.Trajectory(exchanges=TrajectoryExchanges(responses=[exchange])).tokenize( + base_model="base/model", + ).token_ids == [10, 2] + + response = exchange.response.model_dump(mode="python") + response.pop("token_generations", None) + exchange.response = Response.model_validate(response) + with pytest.raises(ValueError, match="no renderable text"): + art.Trajectory(exchanges=TrajectoryExchanges(responses=[exchange])).tokenize( + base_model="base/model", + ) + + +def test_responses_parallel_function_calls_form_one_assistant_turn( + monkeypatch: pytest.MonkeyPatch, +) -> None: + exchange = _response_exchange("parallel-tools", 2) + response_data = exchange.response.model_dump(mode="python") + response_data.pop("token_generations", None) + exchange.response = Response.model_validate(response_data) + exchange.request["input"] = [ + { + "id": "reasoning-1", + "summary": [{"type": "summary_text", "text": "think"}], + "type": "reasoning", + }, + {"type": "function_call", "call_id": "one", "name": "first", "arguments": "{}"}, + { + "type": "function_call", + "call_id": "two", + "name": "second", + "arguments": "{}", + }, + ] + seen: list[list[dict[str, Any]]] = [] + + class Tokenizer(_FakeTokenizer): + def apply_chat_template( + self, messages: list[dict[str, Any]], **kwargs: object + ) -> list[int]: + seen.append(messages) + return super().apply_chat_template(messages, **kwargs) + + monkeypatch.setattr( + "art.trajectories._tokenize._load_tokenizer", lambda _config: Tokenizer() + ) + art.Trajectory(exchanges=TrajectoryExchanges(responses=[exchange])).tokenize( + base_model="base/model", + ) + + assistant = seen[0][0] + assert assistant["reasoning"] == "think" + assert [call["function"]["name"] for call in assistant["tool_calls"]] == [ + "first", + "second", + ] + + +def test_responses_token_generations_preserve_every_generation() -> None: + exchange = _response_exchange("multi-generation", 2) + data = exchange.response.model_dump(mode="python") + data["output"].append( + { + "id": "message-second", + "type": "message", + "role": "assistant", + "status": "completed", + "content": [ + { + "type": "output_text", + "text": "second", + "annotations": [], + "logprobs": [], + } + ], + } + ) + data["token_generations"] = [ + { + "prompt_token_ids": [1], + "output_tokens": [{"token_id": 2, "logprob": -0.2}], + "output_indices": [0], + }, + { + "prompt_token_ids": [1, 2, 3], + "output_tokens": [{"token_id": 4, "logprob": -0.4}], + "output_indices": [1], + }, + ] + exchange.response = Response.model_validate(data) + history = art.Trajectory( + exchanges=TrajectoryExchanges(responses=[exchange]) + ).responses_history() + + tokenized = history.tokenize() + + assert tokenized.token_ids == [1, 2, 3, 4] + assert tokenized.logprobs[1] == -0.2 + assert tokenized.logprobs[3] == -0.4 + assert tokenized.flags == [ + art.TokenFlag.EXACT, + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + art.TokenFlag.EXACT, + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + ] + + +def test_responses_prompt_disagreement_preserves_sampled_token_identity( + monkeypatch: pytest.MonkeyPatch, +) -> None: + exchange = _response_exchange("retokenized-generation", 101) + data = exchange.response.model_dump(mode="python") + data["output"].append( + { + "id": "message-second", + "type": "message", + "role": "assistant", + "status": "completed", + "content": [ + { + "type": "output_text", + "text": "dog", + "annotations": [], + "logprobs": [], + } + ], + } + ) + data["token_generations"] = [ + { + "prompt_token_ids": [1], + "output_tokens": [{"token_id": 101, "logprob": -0.1, "text": "cat"}], + "output_indices": [0], + }, + { + "prompt_token_ids": [1, 500, 3], + "output_tokens": [{"token_id": 4, "logprob": -0.4, "text": "dog"}], + "output_indices": [1], + }, + ] + exchange.response = Response.model_validate(data) + history = art.Trajectory( + exchanges=TrajectoryExchanges(responses=[exchange]) + ).responses_history() + + class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> list[int]: + del kwargs + return {"cat": [500], "dog": [4]}[text] + + def apply_chat_template(self, *args: object, **kwargs: object) -> list[int]: + raise AssertionError("Exact Responses tokenization does not render chat") + + monkeypatch.setattr( + "art.trajectories._tokenize._WARNED_PREFIX_RETOKENIZATION", False + ) + with pytest.warns(UserWarning, match="preserved the original sampled token IDs"): + tokenized = history.tokenize(tokenizer=Tokenizer()) + + assert tokenized.token_ids == [1, 101, 3, 4] + assert tokenized.logprobs[1] == -0.1 + + +@pytest.mark.parametrize( + "token_generations, match", + [ + ([], "must be omitted"), + ( + [ + { + "output_tokens": [{"token_id": 2}], + "output_indices": [0], + } + ], + "prompt_token_ids", + ), + ( + [ + { + "prompt_token_ids": [1], + "output_tokens": [{"token_id": True}], + "output_indices": [0], + } + ], + "exact token ID", + ), + ( + [ + { + "prompt_token_ids": [1], + "output_tokens": [{"token_id": 2}], + "output_indices": [True], + } + ], + "integers", + ), + ( + [ + { + "prompt_token_ids": [1], + "output_tokens": [{"token_id": 2}], + "output_indices": [1], + } + ], + "out of bounds", + ), + ], +) +def test_responses_token_generations_fail_closed( + token_generations: list[dict[str, Any]], match: str +) -> None: + exchange = _response_exchange("invalid-generation", 2) + data = exchange.response.model_dump(mode="python") + data["token_generations"] = token_generations + exchange.response = Response.model_validate(data) + + with pytest.raises(ValueError, match=match): + art.Trajectory( + exchanges=TrajectoryExchanges(responses=[exchange]) + ).responses_history().tokenize() + + +def test_responses_terminal_generation_without_output_items_is_tokenized() -> None: + exchange = _response_exchange("terminal-eos", 2) + data = exchange.response.model_dump(mode="python") + data["output"] = [] + data["token_generations"] = [ + { + "prompt_token_ids": [1], + "output_tokens": [{"token_id": 2, "logprob": -0.2}], + "output_indices": [], + } + ] + exchange.response = Response.model_validate(data) + + history = art.Trajectory( + exchanges=TrajectoryExchanges(responses=[exchange]) + ).responses_history() + tokenized = history.tokenize() + + assert history.input[-1] == {"role": "assistant", "content": ""} + assert history.input_sources[-1] == ResponsesItemSource( + exchange=exchange, generation_index=0 + ) + assert tokenized.token_ids == [1, 2] + assert math.isnan(tokenized.logprobs[0]) + assert tokenized.logprobs[1] == pytest.approx(-0.2) + assert tokenized.flags == [ + art.TokenFlag.EXACT, + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + ] + assert history.as_chat_completions_history().messages[-1] == { + "role": "assistant", + "content": "", + } + + +def test_responses_terminal_generation_without_output_items_survives_rerender() -> None: + exchange = _response_exchange("terminal-eos-rerender", 2) + data = exchange.response.model_dump(mode="python") + data["output"] = [] + data["token_generations"] = [ + { + "prompt_token_ids": [1], + "output_tokens": [{"token_id": 2, "logprob": -0.2}], + "output_indices": [], + } + ] + exchange.response = Response.model_validate(data) + history = ( + art.Trajectory(exchanges=TrajectoryExchanges(responses=[exchange])) + .responses_history() + .as_chat_completions_history() + ) + + class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> list[int]: + del text, kwargs + return [1] + + def apply_chat_template( + self, messages: list[dict[str, Any]], **kwargs: object + ) -> list[int]: + del messages, kwargs + return [1] + + tokenized = history.tokenize(tokenizer=Tokenizer(), chat_template="custom") + + assert tokenized.token_ids == [1, 2] + assert math.isnan(tokenized.logprobs[0]) + assert tokenized.logprobs[1] == pytest.approx(-0.2) + assert tokenized.flags == [ + art.TokenFlag(0), + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + ] + + +def test_responses_empty_chat_source_requires_outputless_generation() -> None: + from art.trajectories._tokenize import _responses_source_generation + + exchange = _response_exchange("nonempty-generation-empty-source", 2) + history = ( + art.Trajectory(exchanges=TrajectoryExchanges(responses=[exchange])) + .responses_history() + .as_chat_completions_history() + ) + assistant_index = next( + index + for index, source in enumerate(history.message_sources) + if source is not None and source.output_indices is not None + ) + history.messages[assistant_index] = {"role": "assistant", "content": ""} + invalid_source = ChatCompletionsMessageSource( + exchange=exchange, + output_indices=(), + generation_index=0, + ) + history.message_sources[assistant_index] = invalid_source + + with pytest.raises(ValueError, match="empty output source"): + _responses_source_generation(invalid_source) + with pytest.raises(ValueError, match="empty output source"): + history.tokenize() + + +def test_responses_request_composite_source_validation_uses_contiguous_items() -> None: + from art.trajectories._tokenize import _validate_history_sources + + exchange = _response_exchange("request-composite", 2) + exchange.request["input"] = [ + { + "type": "function_call", + "call_id": f"call-{index}", + "name": f"tool_{index}", + "arguments": "{}", + } + for index in range(2) + ] + history = ( + art.Trajectory(exchanges=TrajectoryExchanges(responses=[exchange])) + .responses_history() + .as_chat_completions_history() + ) + + assert len(history.messages[0].get("tool_calls", [])) == 2 + _validate_history_sources(history) + + +def test_responses_nonterminal_generation_without_output_items_raises() -> None: + exchange = _response_exchange("hidden-control", 2) + data = exchange.response.model_dump(mode="python") + data["token_generations"] = [ + { + "prompt_token_ids": [1], + "output_tokens": [{"token_id": 2, "logprob": -0.2}], + "output_indices": [], + }, + { + "prompt_token_ids": [1, 2], + "output_tokens": [{"token_id": 3, "logprob": -0.3}], + "output_indices": [0], + }, + ] + exchange.response = Response.model_validate(data) + + with pytest.raises(ValueError, match="nonterminal"): + art.Trajectory( + exchanges=TrajectoryExchanges(responses=[exchange]) + ).responses_history() + + +def test_tokenization_rejects_mutated_mixed_representation() -> None: + trajectory = art.Trajectory( + messages_and_choices=[{"role": "user", "content": "hi"}] + ) + trajectory.exchanges.chat_completions.append(_chat_exchange([1], [2])) + + with pytest.raises(ValueError, match="both exchanges and legacy histories"): + trajectory.tokenize() + + +def test_responses_previous_response_id_resolves_local_history( + monkeypatch: pytest.MonkeyPatch, +) -> None: + class Tokenizer: + def apply_chat_template( + self, messages: list[dict[str, Any]], **kwargs: object + ) -> list[int]: + del kwargs + return [10] if len(messages) == 1 else [10, 20, 11] + + monkeypatch.setattr( + "art.trajectories._tokenize._load_tokenizer", lambda _config: Tokenizer() + ) + first = _response_exchange("resp-1", 20, prompt_token_ids=[10]) + second = _response_exchange( + "resp-2", + 30, + previous_response_id="resp-1", + offset=1, + prompt_token_ids=[10, 20, 11], + ) + trajectory = art.Trajectory( + exchanges=TrajectoryExchanges(responses=[first, second]) + ) + + assert trajectory.tokenize(base_model="base/model").token_ids == [ + 10, + 20, + 11, + 30, + ] + + second.request["previous_response_id"] = "missing" + assert len(trajectory.responses_histories()) == 2 + with pytest.raises(ValueError, match="exactly one history"): + trajectory.tokenize(base_model="base/model") + assert trajectory.responses_histories()[1].tokenize( + base_model="base/model" + ).token_ids == [10, 20, 11, 30] + + +def test_prefix_retokenization_preserves_sampled_ids_and_logprobs( + monkeypatch: pytest.MonkeyPatch, +) -> None: + first = _chat_exchange([1], [101, 102]) + first.response.choices[0].message.content = "cat" + second = _chat_exchange([1, 500, 3], [4], offset=1) + second.request["messages"] = [ + {"role": "user", "content": "turn 0"}, + {"role": "assistant", "content": "cat"}, + {"role": "user", "content": "turn 1"}, + ] + + class Tokenizer: + def __call__(self, text: str, *, add_special_tokens: bool = False) -> list[int]: + assert not add_special_tokens + return {"cat": [500], "answer": [4], "turn 0": [1], "turn 1": [3]}[text] + + def apply_chat_template( + self, messages: list[dict[str, Any]], **kwargs: object + ) -> list[int]: + del kwargs + rendered: list[int] = [] + for message in messages: + rendered.extend(self(str(message["content"]))) + return rendered + + monkeypatch.setattr( + "art.trajectories._tokenize._load_tokenizer", lambda _config: Tokenizer() + ) + monkeypatch.setattr( + "art.trajectories._tokenize._WARNED_PREFIX_RETOKENIZATION", False + ) + trajectory = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[first, second]) + ) + + with pytest.warns(UserWarning, match="preserved the original sampled token IDs"): + tokenized = trajectory.tokenize(base_model="base/model") + + assert tokenized.token_ids == [1, 101, 102, 3, 4] + assert tokenized.logprobs[1:3] == [-10.1, -10.2] + assert all(tokenized.flags[index] & art.TokenFlag.EXACT for index in (1, 2, 4)) + + +def test_template_change_rerenders_scaffold_but_preserves_sampled_output() -> None: + history = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[_chat_exchange([1], [2])]) + ).chat_completions_history() + history.chat_template = "custom" + + class Tokenizer: + def __call__(self, text: str, *, add_special_tokens: bool = False) -> list[int]: + assert not add_special_tokens + return {"turn 0": [10], "answer": [20]}[text] + + def apply_chat_template( + self, + messages: list[dict[str, Any]], + *, + tools: object, + tokenize: bool, + add_generation_prompt: bool, + chat_template: str | None = None, + **kwargs: object, + ) -> list[int]: + del tools, tokenize, add_generation_prompt, kwargs + assert chat_template == "custom" + if len(messages) == 1: + return [10] + if str(messages[-1]["content"]).startswith("ART_TRAJECTORY_"): + return [10, 999, 30] + return [10, 20, 30] + + tokenized = history.tokenize(tokenizer=Tokenizer()) + + assert tokenized.token_ids == [10, 2, 30] + assert tokenized.logprobs[1] == -0.2 + assert tokenized.flags[1] == art.TokenFlag.EXACT | art.TokenFlag.SAMPLED + + +def test_template_change_preserves_complete_exact_sampled_suffix() -> None: + exchange = _chat_exchange([1], [2, 3]) + history = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[exchange]) + ).chat_completions_history() + history.chat_template = "custom" + + class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> list[int]: + del kwargs + return {"turn 0": [1], "answer": [2]}[text] + + def apply_chat_template( + self, messages: list[dict[str, Any]], **kwargs: object + ) -> list[int]: + del kwargs + if messages[-1]["role"] != "assistant": + return [1] + if str(messages[-1]["content"]).startswith("ART_TRAJECTORY_"): + return [1, 999, 9] + return [1, 2, 9] + + tokenized = history.tokenize(tokenizer=Tokenizer()) + + assert tokenized.token_ids == [1, 2, 3, 9] + assert tokenized.logprobs[1:3] == pytest.approx([-0.2, -0.3]) + assert math.isnan(tokenized.logprobs[3]) + assert tokenized.flags[1:] == [ + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + art.TokenFlag(0), + ] + + +def test_responses_generation_evidence_is_atomic_and_partial_edits_do_not_replay() -> ( + None +): + exchange = _response_exchange("reasoning-and-answer", 0) + data = exchange.response.model_dump(mode="python") + data["output"] = [ + { + "id": "reasoning", + "type": "reasoning", + "summary": [{"type": "summary_text", "text": "think"}], + }, + { + "id": "message", + "type": "message", + "role": "assistant", + "status": "completed", + "content": [ + { + "type": "output_text", + "text": "answer", + "annotations": [], + "logprobs": [], + } + ], + }, + ] + data["token_generations"] = [ + { + "prompt_token_ids": [1], + "output_tokens": [ + {"token_id": 2, "logprob": -0.2, "text": "think"}, + {"token_id": 3, "logprob": -0.3, "text": "answer"}, + ], + "output_indices": [0, 1], + } + ] + exchange.response = Response.model_validate(data) + + class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> object: + if kwargs.get("return_offsets_mapping"): + answer_start = text.index("answer") + answer_end = answer_start + len("answer") + offsets: list[tuple[int, int]] + token_ids = [1] + if "think" in text: + think_start = text.index("think") + think_end = think_start + len("think") + offsets = [ + (0, think_start), + (think_start, think_end), + (answer_start, answer_end), + (answer_end, len(text)), + ] + token_ids.extend([20, 30, 9]) + else: + offsets = [ + (0, answer_start), + (answer_start, answer_end), + (answer_end, len(text)), + ] + token_ids.extend([30, 9]) + return {"input_ids": token_ids, "offset_mapping": offsets} + return {"turn 0": [1], "think": [20], "answer": [30]}[text] + + def apply_chat_template( + self, messages: list[dict[str, Any]], **kwargs: object + ) -> object: + tokenize = kwargs.pop("tokenize") + del kwargs + if messages[-1]["role"] != "assistant": + return [1] + assistant = messages[-1] + rendered = ( + f"{messages[0]['content']}" + + ( + f"{assistant['reasoning']}" + if assistant.get("reasoning") + else "" + ) + + f"{assistant['content']}" + ) + if not tokenize: + return rendered + if assistant.get("reasoning"): + return [1, 20, 30, 9] + return [1, 30, 9] + + history = art.Trajectory( + exchanges=TrajectoryExchanges(responses=[exchange]) + ).responses_history() + exact = history.tokenize(tokenizer=Tokenizer(), chat_template="custom") + assert exact.token_ids == [1, 2, 3, 9] + assert exact.logprobs[1:3] == pytest.approx([-0.2, -0.3]) + assert math.isnan(exact.logprobs[3]) + + del history.input[1] + del history.input_sources[1] + partial = history.tokenize(tokenizer=Tokenizer(), chat_template="custom") + assert partial.token_ids == [1, 30, 9] + assert 2 not in partial.token_ids + assert partial.flags[1] == art.TokenFlag.SAMPLED + assert math.isnan(partial.logprobs[1]) + + +def test_responses_generation_source_rejects_content_from_another_generation() -> None: + exchange = _response_exchange("generation-provenance", 2) + data = exchange.response.model_dump(mode="python") + data["output"].append( + { + "id": "message-second-generation", + "type": "message", + "role": "assistant", + "status": "completed", + "content": [ + { + "type": "output_text", + "text": "second", + "annotations": [], + "logprobs": [], + } + ], + } + ) + data["token_generations"] = [ + { + "prompt_token_ids": [1], + "output_tokens": [{"token_id": 2, "logprob": -0.2}], + "output_indices": [0], + }, + { + "prompt_token_ids": [1, 2, 3], + "output_tokens": [{"token_id": 4, "logprob": -0.4}], + "output_indices": [1], + }, + ] + exchange.response = Response.model_validate(data) + history = ( + art.Trajectory(exchanges=TrajectoryExchanges(responses=[exchange])) + .responses_history() + .as_chat_completions_history() + ) + first_generation = next( + index + for index, source in enumerate(history.message_sources) + if source is not None + and source.generation_index == 0 + and history.messages[index].get("role") == "assistant" + ) + history.messages[first_generation] = { + "role": "assistant", + "content": "second", + } + + with pytest.raises(ValueError, match="no longer matches its source exchange"): + history.tokenize() + + +def test_responses_chat_rerender_preserves_equal_length_generation_evidence() -> None: + exchange = _response_exchange("equal-length-generations", 2) + data = exchange.response.model_dump(mode="python") + data["output"].append( + { + "id": "message-second-generation", + "type": "message", + "role": "assistant", + "status": "completed", + "content": [ + { + "type": "output_text", + "text": "second", + "annotations": [], + "logprobs": [], + } + ], + } + ) + data["token_generations"] = [ + { + "prompt_token_ids": [1], + "output_tokens": [{"token_id": 2, "logprob": -0.2}], + "output_indices": [0], + }, + { + "prompt_token_ids": [1, 2, 3], + "output_tokens": [{"token_id": 4, "logprob": -0.4}], + "output_indices": [1], + }, + ] + exchange.response = Response.model_validate(data) + history = ( + art.Trajectory(exchanges=TrajectoryExchanges(responses=[exchange])) + .responses_history() + .as_chat_completions_history() + ) + + class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> object: + mapping = {"turn 0": 10, "answer": 20, "second": 40} + if not kwargs.get("return_offsets_mapping"): + return [mapping[text]] + token_ids: list[int] = [] + offsets: list[tuple[int, int]] = [] + for match in re.finditer(r"(.*?)|", text): + if match.group(1) is None: + token_ids.append(30) + offsets.append(match.span()) + else: + token_ids.append(mapping[match.group(1)]) + offsets.append(match.span(1)) + return {"input_ids": token_ids, "offset_mapping": offsets} + + def apply_chat_template( + self, + messages: list[dict[str, Any]], + *, + add_generation_prompt: bool, + tokenize: bool, + **kwargs: object, + ) -> object: + del kwargs + result = "" + for index, message in enumerate(messages): + result += f"{message['content']}" + if index == 1 and (len(messages) > 2 or add_generation_prompt): + result += "" + return self(result, return_offsets_mapping=True) if tokenize else result + + tokenized = history.tokenize(tokenizer=Tokenizer(), chat_template="custom") + + assert tokenized.token_ids == [10, 2, 30, 4] + assert 20 not in tokenized.token_ids + assert 40 not in tokenized.token_ids + assert tokenized.logprobs[1::2] == pytest.approx([-0.2, -0.4]) + assert tokenized.flags == [ + art.TokenFlag(0), + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + art.TokenFlag(0), + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + ] + + +def test_responses_chat_sampled_source_requires_generation_identity() -> None: + exchange = _response_exchange("missing-generation-identity", 2) + history = ( + art.Trajectory(exchanges=TrajectoryExchanges(responses=[exchange])) + .responses_history() + .as_chat_completions_history() + ) + assistant_index = next( + index + for index, source in enumerate(history.message_sources) + if source is not None and source.output_indices is not None + ) + source = history.message_sources[assistant_index] + assert source is not None + history.message_sources[assistant_index] = ChatCompletionsMessageSource( + exchange=source.exchange, + output_indices=source.output_indices, + ) + + class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> list[int]: + del kwargs + return {"turn 0": [1], "answer": [2]}[text] + + def apply_chat_template( + self, messages: list[dict[str, Any]], **kwargs: object + ) -> list[int]: + del kwargs + return [ + token for message in messages for token in self(str(message["content"])) + ] + + with pytest.raises(ValueError, match="no generation identity"): + history.tokenize(tokenizer=Tokenizer()) + + +def test_responses_chat_output_indices_reuse_one_complete_generation() -> None: + exchange = _response_exchange( + "response-output-indices", + 2, + prompt_token_ids=[1], + ) + history = art.ChatCompletionsHistory( + model="test/model", + messages=[ + {"role": "user", "content": "turn 0"}, + {"role": "assistant", "content": "answer"}, + ], + message_sources=[ + ChatCompletionsMessageSource(exchange=exchange, request_index=0), + ChatCompletionsMessageSource( + exchange=exchange, + output_indices=(0,), + generation_index=0, + ), + ], + ) + + tokenized = history.tokenize() + + assert tokenized.token_ids == [1, 2] + assert tokenized.logprobs[-1] == pytest.approx(-0.1) + assert tokenized.flags[-1] == art.TokenFlag.EXACT | art.TokenFlag.SAMPLED + + +def test_responses_chat_output_indices_are_bounds_checked() -> None: + exchange = _response_exchange("response-output-indices-invalid", 2) + history = art.ChatCompletionsHistory( + model="test/model", + messages=[{"role": "assistant", "content": "answer"}], + message_sources=[ + ChatCompletionsMessageSource( + exchange=exchange, + output_indices=(1,), + generation_index=0, + ) + ], + ) + + with pytest.raises(ValueError, match="out of bounds"): + history.tokenize() + + +def _multi_output_responses_chat_history() -> art.ChatCompletionsHistory: + exchange = _response_exchange("multi-output-generation", 2) + data = exchange.response.model_dump(mode="python") + data["output"][0]["content"][0]["text"] = "first" + data["output"].append( + { + "id": "message-second-output", + "type": "message", + "role": "assistant", + "status": "completed", + "content": [ + { + "type": "output_text", + "text": "second", + "annotations": [], + "logprobs": [], + } + ], + } + ) + data["token_generations"] = [ + { + "prompt_token_ids": [1], + "output_tokens": [ + {"token_id": 2, "logprob": -0.2}, + {"token_id": 3, "logprob": -0.3}, + ], + "output_indices": [0, 1], + } + ] + exchange.response = Response.model_validate(data) + return ( + art.Trajectory(exchanges=TrajectoryExchanges(responses=[exchange])) + .responses_history() + .as_chat_completions_history() + ) + + +def test_responses_output_source_rejects_sibling_generation_output() -> None: + history = _multi_output_responses_chat_history() + first_output = next( + index + for index, source in enumerate(history.message_sources) + if source is not None and source.output_indices == (0,) + ) + history.messages[first_output] = { + "role": "assistant", + "content": "second", + } + + with pytest.raises(ValueError, match="no longer matches its source exchange"): + history.tokenize() + + +def test_responses_multi_output_generation_is_rendered_without_duplicate_evidence() -> ( + None +): + history = _multi_output_responses_chat_history() + + class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> list[int]: + del kwargs + return {"turn 0": [1], "first": [20], "second": [30]}[text] + + def apply_chat_template( + self, messages: list[dict[str, Any]], **kwargs: object + ) -> list[int]: + del kwargs + result: list[int] = [] + for message in messages: + result.extend(self(str(message["content"]))) + return result + + tokenized = history.tokenize(tokenizer=Tokenizer(), chat_template="custom") + + assert tokenized.token_ids == [1, 20, 30] + assert all(math.isnan(value) for value in tokenized.logprobs) + assert tokenized.flags == [ + art.TokenFlag(0), + art.TokenFlag.SAMPLED, + art.TokenFlag.SAMPLED, + ] + + +def test_responses_multi_output_chat_conversion_preserves_item_logprobs() -> None: + projected = _multi_output_responses_chat_history() + exchange = next( + source.exchange + for source in projected.message_sources + if source is not None and source.output_indices == (0,) + ) + assert isinstance(exchange, ResponsesExchange) + data = exchange.response.model_dump(mode="python") + data.pop("token_generations") + for output, (text, logprob) in zip( + data["output"], (("first", -0.1), ("second", -0.2)), strict=True + ): + output["content"][0]["logprobs"] = [ + { + "token": text, + "logprob": logprob, + "bytes": list(text.encode()), + "top_logprobs": [], + } + ] + exchange.response = Response.model_validate(data) + history = ( + art.Trajectory(exchanges=TrajectoryExchanges(responses=[exchange])) + .responses_history() + .as_chat_completions_history() + ) + + class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> list[int]: + del kwargs + return {"turn 0": [1], "first": [2], "second": [3]}[text] + + def apply_chat_template( + self, messages: list[dict[str, Any]], **kwargs: object + ) -> list[int]: + del kwargs + return [ + token for message in messages for token in self(str(message["content"])) + ] + + tokenized = history.tokenize(tokenizer=Tokenizer()) + + assert tokenized.token_ids == [1, 2, 3] + assert math.isnan(tokenized.logprobs[0]) + assert tokenized.logprobs[1:] == [-0.1, -0.2] + assert tokenized.flags == [ + art.TokenFlag(0), + art.TokenFlag.SAMPLED, + art.TokenFlag.SAMPLED, + ] + + +def test_mutable_chat_history_is_authoritative_and_does_not_replay_removed_turns() -> ( + None +): + first = _chat_exchange([10], [20]) + first.request["messages"] = [{"role": "user", "content": "first"}] + second = _chat_exchange([10, 20, 30], [40], offset=1) + second.request["messages"] = [ + {"role": "user", "content": "first"}, + {"role": "assistant", "content": "answer"}, + {"role": "user", "content": "second"}, + ] + history = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[first, second]) + ).chat_completions_history() + del history.messages[:2] + del history.message_sources[:2] + + class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> list[int]: + del kwargs + return {"second": [31], "answer": [41]}[text] + + def apply_chat_template( + self, + messages: list[dict[str, Any]], + *, + add_generation_prompt: bool, + **kwargs: object, + ) -> list[int]: + del kwargs + contents = [message["content"] for message in messages] + if contents == ["second", "answer"]: + return [30, 41, 50] + if len(contents) == 2 and str(contents[-1]).startswith("ART_TRAJECTORY_"): + return [30, 99, 50] + assert contents == ["second"] + return [30] if add_generation_prompt else [30] + + tokenized = history.tokenize(tokenizer=Tokenizer()) + + assert tokenized.token_ids == [30, 40, 50] + assert 20 not in tokenized.token_ids + assert tokenized.flags == [ + art.TokenFlag(0), + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + art.TokenFlag(0), + ] + + +def test_request_assistant_messages_are_not_marked_sampled() -> None: + exchange = _chat_exchange([10], [40]) + exchange.request["messages"] = [ + {"role": "assistant", "content": "seed"}, + {"role": "user", "content": "question"}, + ] + history = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[exchange]) + ).chat_completions_history() + history.chat_template = "rerender" + + class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> list[int]: + del kwargs + return {"seed": [20], "question": [30], "answer": [41]}[text] + + def apply_chat_template( + self, + messages: list[dict[str, Any]], + *, + add_generation_prompt: bool, + **kwargs: object, + ) -> list[int]: + del kwargs + if messages[-1]["role"] == "assistant" and len(messages) == 3: + if str(messages[-1]["content"]).startswith("ART_TRAJECTORY_"): + return [5, 20, 6, 30, 7, 99, 8] + return [5, 20, 6, 30, 7, 41, 8] + assert add_generation_prompt + return [5, 20, 6, 30, 7] + + tokenized = history.tokenize(tokenizer=Tokenizer()) + + assert tokenized.token_ids == [5, 20, 6, 30, 7, 40, 8] + assert not tokenized.flags[1] & art.TokenFlag.SAMPLED + assert tokenized.flags[5] == art.TokenFlag.EXACT | art.TokenFlag.SAMPLED + + +def test_rerender_constrains_exact_output_to_its_message_region() -> None: + exchange = _chat_exchange([7], [7]) + exchange.request["messages"] = [{"role": "user", "content": "same"}] + exchange.response.choices[0].message.content = "same" + history = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[exchange]) + ).chat_completions_history() + history.chat_template = "rerender" + + class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> list[int]: + del kwargs + assert text == "same" + return [7] + + def apply_chat_template( + self, + messages: list[dict[str, Any]], + *, + add_generation_prompt: bool, + **kwargs: object, + ) -> list[int]: + del kwargs + if messages[-1]["role"] == "assistant": + return [7, 99, 7] + assert add_generation_prompt + return [7, 99] + + tokenized = history.tokenize(tokenizer=Tokenizer()) + + assert tokenized.flags == [ + art.TokenFlag(0), + art.TokenFlag(0), + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + ] + assert math.isnan(tokenized.logprobs[0]) + assert tokenized.logprobs[2] == -0.7 + + +def test_rerender_does_not_bind_sampled_ids_to_token_equivalent_user_text() -> None: + exchange = _chat_exchange([7], [7]) + exchange.request["messages"] = [{"role": "user", "content": "cat"}] + exchange.response.choices[0].message.content = "dog" + history = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[exchange]) + ).chat_completions_history() + history.chat_template = "rerender" + + class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> list[int]: + del text, kwargs + return [7] + + def apply_chat_template( + self, messages: list[dict[str, Any]], **kwargs: object + ) -> list[int]: + del messages, kwargs + return [7, 99, 7] + + with pytest.raises(ValueError, match="uniquely locate"): + history.tokenize(tokenizer=Tokenizer()) + + +def test_rerender_does_not_bind_unique_exact_id_outside_sampled_message() -> None: + exchange = _chat_exchange([42], [7]) + exchange.request["messages"] = [{"role": "user", "content": "cat"}] + exchange.response.choices[0].message.content = "dog" + history = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[exchange]) + ).chat_completions_history() + history.chat_template = "rerender" + + class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> list[int]: + del kwargs + return {"cat": [7], "dog": [500]}[text] + + def apply_chat_template( + self, + messages: list[dict[str, Any]], + *, + add_generation_prompt: bool, + **kwargs: object, + ) -> list[int]: + del kwargs + if messages[-1]["role"] == "assistant": + return [42, 7, 99, 500] + return [42, 7, 99] if add_generation_prompt else [42, 7] + + tokenized = history.tokenize(tokenizer=Tokenizer()) + + assert tokenized.token_ids == [42, 7, 99, 7] + assert not tokenized.flags[1] & art.TokenFlag.SAMPLED + assert tokenized.flags[-1] == (art.TokenFlag.EXACT | art.TokenFlag.SAMPLED) + assert math.isnan(tokenized.logprobs[1]) + assert tokenized.logprobs[-1] == -0.7 + + +def test_rerender_rejects_sampled_text_ambiguous_with_trailing_scaffold() -> None: + exchange = _chat_exchange([42], [7]) + exchange.request["messages"] = [{"role": "user", "content": "cat"}] + exchange.response.choices[0].message.content = "dog" + history = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[exchange]) + ).chat_completions_history() + history.chat_template = "rerender" + + class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> list[int]: + del kwargs + return {"cat": [1], "dog": [500]}[text] + + def apply_chat_template( + self, messages: list[dict[str, Any]], **kwargs: object + ) -> list[int]: + del messages, kwargs + return [100, 500, 99, 500] + + with pytest.raises(ValueError, match="uniquely locate"): + history.tokenize(tokenizer=Tokenizer()) + + +def test_rerender_does_not_duplicate_sampled_trailing_eos() -> None: + exchange = _chat_exchange([1], [7, 2]) + exchange.request["messages"] = [{"role": "user", "content": "question"}] + exchange.response.choices[0].message.content = "answer" + history = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[exchange]) + ).chat_completions_history() + history.chat_template = "rerender" + + class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> list[int]: + del kwargs + return {"question": [1], "answer": [7]}[text] + + def apply_chat_template( + self, + messages: list[dict[str, Any]], + *, + add_generation_prompt: bool, + **kwargs: object, + ) -> list[int]: + del kwargs + if messages[-1]["role"] == "assistant": + return [1, 99, 7, 2] + assert add_generation_prompt + return [1, 99] + + tokenized = history.tokenize(tokenizer=Tokenizer()) + + assert tokenized.token_ids == [1, 99, 7, 2] + assert tokenized.token_ids.count(2) == 1 + assert tokenized.flags[-2:] == [ + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + ] + assert tokenized.logprobs[-2:] == [-0.7, -0.2] + + +def test_chat_view_preserves_initial_prompt_and_ignores_later_disagreement() -> None: + first = _chat_exchange([1], [2]) + second = _chat_exchange([9, 8, 7], [3], offset=1) + trajectory = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[first, second]) + ) + + class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> list[int]: + del kwargs + return {"turn 0": [100], "answer": [2], "turn 1": [7]}[text] + + def apply_chat_template( + self, messages: list[dict[str, Any]], **kwargs: object + ) -> list[int]: + del kwargs + return [ + token for message in messages for token in self(str(message["content"])) + ] + + tokenized = trajectory.tokenize(tokenizer=Tokenizer()) + + assert tokenized.token_ids == [1, 2, 7, 3] + + +def test_reasoning_stripped_chat_histories_tokenize_authoritative_views() -> None: + first = _chat_exchange([1], [2, 101, 102, 9]) + first.request["messages"] = [{"role": "user", "content": "one"}] + first_data = first.response.model_dump(mode="python") + first_data["choices"][0]["message"] = { + "role": "assistant", + "content": "first", + "reasoning": "thought-one", + } + first.response = ChatCompletion.model_validate(first_data) + + second = _chat_exchange([1, 101, 102, 9, 4], [5, 6], offset=1) + second.request["messages"] = [ + {"role": "user", "content": "one"}, + {"role": "assistant", "content": "first"}, + {"role": "user", "content": "two"}, + ] + second_data = second.response.model_dump(mode="python") + second_data["choices"][0]["message"] = { + "role": "assistant", + "content": "second", + "reasoning": "thought-two", + } + second.response = ChatCompletion.model_validate(second_data) + trajectory = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[first, second]) + ) + + class Tokenizer: + name_or_path = "test/model" + + def __call__(self, text: str, **kwargs: object) -> list[int]: + del kwargs + return { + "one": [1], + "first": [500], + "two": [4], + "thought-two": [5], + "second": [6], + }[text] + + def apply_chat_template( + self, messages: list[dict[str, Any]], **kwargs: object + ) -> list[int]: + del kwargs + assert [message.get("content") for message in messages] == [ + "one", + "first", + "two", + "second", + ] + return [1, 500, 9, 4, 5, 6] + + tokenized = trajectory.tokenize(multi_history=True, tokenizer=Tokenizer()) + + assert [history.token_ids for history in tokenized.histories] == [ + [1, 2, 101, 102, 9], + [1, 101, 102, 9, 4, 5, 6], + ] + assert tokenized.histories[1].flags[1] & art.TokenFlag.SAMPLED + assert tokenized.histories[1].flags[1] & art.TokenFlag.EXACT + assert tokenized.histories[1].logprobs[1:3] == [-10.1, -10.2] + assert tokenized.histories[1].flags[3] == ( + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED + ) + assert tokenized.histories[1].logprobs[3] == -0.9 + assert 2 not in tokenized.histories[1].token_ids + assert 500 not in tokenized.histories[1].token_ids + + +def test_reasoning_stripped_histories_remain_trainable_end_to_end( + monkeypatch: pytest.MonkeyPatch, +) -> None: + first = _chat_exchange([1], [2, 101, 102, 9]) + first.request["messages"] = [{"role": "user", "content": "one"}] + first_data = first.response.model_dump(mode="python") + first_data["choices"][0]["message"] = { + "role": "assistant", + "content": "first", + "reasoning": "thought-one", + } + first.response = ChatCompletion.model_validate(first_data) + second = _chat_exchange([1, 101, 102, 9, 4], [5, 6], offset=1) + second.request["messages"] = [ + {"role": "user", "content": "one"}, + {"role": "assistant", "content": "first"}, + {"role": "user", "content": "two"}, + ] + second_data = second.response.model_dump(mode="python") + second_data["choices"][0]["message"] = { + "role": "assistant", + "content": "second", + "reasoning": "thought-two", + } + second.response = ChatCompletion.model_validate(second_data) + + class Tokenizer: + name_or_path = "test/model" + + def __call__(self, text: str, **kwargs: object) -> list[int]: + del kwargs + return { + "one": [1], + "first": [500], + "two": [4], + "thought-two": [5], + "second": [6], + }[text] + + def apply_chat_template( + self, messages: list[dict[str, Any]], **kwargs: object + ) -> list[int]: + del kwargs + content = [message.get("content") for message in messages] + return ( + [1, 2, 101, 102, 9] + if content == ["one", "first"] + else [1, 500, 9, 4, 5, 6] + ) + + monkeypatch.setattr( + "art.trajectories._tokenize._load_tokenizer", lambda _: Tokenizer() + ) + trajectories = [ + art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[first, second]), + reward=reward, + ) + for reward in (1.0, 0.0) + ] + group = art.TrajectoryGroup(trajectories=trajectories) + + from art.preprocessing.tokenize import tokenize_trajectory_groups + from art.tinker_native.data import trajectory_groups_to_datums + + preprocessing = list( + tokenize_trajectory_groups( + Tokenizer(), # type: ignore[arg-type, ty:invalid-argument-type] + [group], + allow_training_without_logprobs=False, + scale_rewards=False, + shuffle_group_trajectories=False, + drop_zero_advantage_trajectories=False, + ) + ) + datums = trajectory_groups_to_datums( + [group], + renderer=None, + tokenizer=None, + normalize_advantages=False, + ) + + assert len(preprocessing) == 4 + assert len(datums) == 4 + assert all( + not math.isnan(logprob) + for result in preprocessing + for logprob, sampled in zip(result.logprobs, result.assistant_mask, strict=True) + if sampled + ) + + +def test_reasoning_stripped_tool_call_keeps_exact_evidence_for_strict_training( + monkeypatch: pytest.MonkeyPatch, +) -> None: + first = _chat_exchange([1], [2, 7, 8]) + first.request["messages"] = [{"role": "user", "content": "one"}] + first_data = first.response.model_dump(mode="python") + first_data["choices"][0]["message"] = { + "role": "assistant", + "content": None, + "reasoning": "thought", + "tool_calls": [ + { + "id": "call-1", + "type": "function", + "function": {"name": "lookup", "arguments": "{}"}, + } + ], + } + first.response = ChatCompletion.model_validate(first_data) + second = _chat_exchange([1, 7, 8, 4], [5], offset=1) + second.request["messages"] = [ + {"role": "user", "content": "one"}, { - "system": [{"type": "text", "text": "system"}], - "messages": [ - {"role": "user", "content": "question"}, - { - "role": "assistant", - "content": [ - {"type": "thinking", "thinking": "reason"}, - {"type": "text", "text": "calling"}, - { - "type": "tool_use", - "id": "call-1", - "name": "lookup", - "input": {"key": "value"}, - }, - ], - }, - { - "role": "user", - "content": [ - { - "type": "tool_result", - "tool_use_id": "call-1", - "content": [{"type": "text", "text": "result"}], - }, - {"type": "text", "text": "continue"}, - ], - }, - ], - } + "role": "assistant", + "content": None, + "tool_calls": first_data["choices"][0]["message"]["tool_calls"], + }, + {"role": "user", "content": "two"}, + ] + + class Tokenizer: + name_or_path = "test/model" + + def __call__(self, text: str, **kwargs: object) -> list[int]: + del kwargs + return {"lookup": [7], "{}": [8]}[text] + + def apply_chat_template( + self, messages: list[dict[str, Any]], **kwargs: object + ) -> list[int]: + del kwargs + assert len(messages) == 4 + return [1, 7, 8, 4, 5] + + monkeypatch.setattr( + "art.trajectories._tokenize._load_tokenizer", lambda _: Tokenizer() ) + trajectories = [ + art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[first, second]), + reward=reward, + ) + for reward in (1.0, 0.0) + ] + group = art.TrajectoryGroup(trajectories=trajectories) - assert messages == [ - {"role": "system", "content": "system"}, - {"role": "user", "content": "question"}, + from art.preprocessing.tokenize import tokenize_trajectory_groups + from art.tinker_native.data import trajectory_groups_to_datums + + tokenized = trajectories[0].tokenize( + multi_history=True, + tokenizer=Tokenizer(), + ) + second_history = tokenized.histories[1] + assert second_history.token_ids == [1, 7, 8, 4, 5] + assert second_history.logprobs[1:3] == [-0.7, -0.8] + assert second_history.flags[1:3] == [ + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + ] + + preprocessing = list( + tokenize_trajectory_groups( + Tokenizer(), # type: ignore[arg-type, ty:invalid-argument-type] + [group], + allow_training_without_logprobs=False, + scale_rewards=False, + shuffle_group_trajectories=False, + drop_zero_advantage_trajectories=False, + ) + ) + datums = trajectory_groups_to_datums( + [group], + renderer=None, + tokenizer=None, + normalize_advantages=False, + ) + + assert len(preprocessing) == 4 + assert len(datums) == 4 + + +def test_responses_prompt_repair_uses_native_text_and_source_position() -> None: + exchange = _response_exchange("repeated-retokenization", 101) + data = exchange.response.model_dump(mode="python") + data["output"].append( { + "id": "message-second", + "type": "message", "role": "assistant", - "content": "calling", - "reasoning": "reason", - "tool_calls": [ + "status": "completed", + "content": [ { - "id": "call-1", - "type": "function", - "function": { - "name": "lookup", - "arguments": '{"key": "value"}', - }, + "type": "output_text", + "text": "dog", + "annotations": [], + "logprobs": [], } ], + } + ) + data["token_generations"] = [ + { + "prompt_token_ids": [500], + "output_tokens": [{"token_id": 101, "logprob": -0.1}], + "output_indices": [0], + }, + { + "prompt_token_ids": [500, 500, 3], + "output_tokens": [{"token_id": 4, "logprob": -0.4}], + "output_indices": [1], }, - {"role": "tool", "tool_call_id": "call-1", "content": "result"}, - {"role": "user", "content": "continue"}, ] + exchange.response = Response.model_validate(data) + history = art.Trajectory( + exchanges=TrajectoryExchanges(responses=[exchange]) + ).responses_history() + + class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> list[int]: + del kwargs + return {"answer": [500], "dog": [4]}[text] + def apply_chat_template(self, *args: object, **kwargs: object) -> list[int]: + raise AssertionError("Exact Responses tokenization must not render chat") -def test_choice_logprobs_survive_tokenizer_fallback( - monkeypatch: pytest.MonkeyPatch, + with pytest.warns(UserWarning, match="preserved the original sampled token IDs"): + tokenized = history.tokenize(tokenizer=Tokenizer()) + + assert tokenized.token_ids == [500, 101, 3, 4] + assert tokenized.logprobs[1] == -0.1 + + +def test_rerender_marks_tool_call_only_generated_region_sampled() -> None: + exchange = _chat_exchange([1], [2]) + data = exchange.response.model_dump(mode="python") + choice = data["choices"][0] + choice.pop("token_ids") + choice["logprobs"] = None + choice["message"] = { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call-1", + "type": "function", + "function": {"name": "lookup", "arguments": '{"x":1}'}, + } + ], + } + exchange.response = ChatCompletion.model_validate(data) + history = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[exchange]) + ).chat_completions_history() + history.chat_template = "rerender" + + class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> object: + if kwargs.get("return_offsets_mapping"): + name_start = text.index("lookup") + name_end = name_start + len("lookup") + args_start = text.index('{"x":1}') + args_end = args_start + len('{"x":1}') + midpoint = args_start + 3 + return { + "input_ids": [1, 10, 20, 25, 30, 31, 26], + "offset_mapping": [ + (0, text.index("")), + (text.index(""), name_start), + (name_start, name_end), + (name_end, args_start), + (args_start, midpoint), + (midpoint, args_end), + (args_end, len(text)), + ], + } + return { + "turn 0": [1], + "lookup": [20], + '{"x":1}': [30, 31], + }[text] + + def apply_chat_template( + self, + messages: list[dict[str, Any]], + *, + add_generation_prompt: bool, + **kwargs: object, + ) -> object: + tokenize = kwargs.pop("tokenize") + del kwargs + if messages[-1]["role"] == "assistant": + function = messages[-1]["tool_calls"][0]["function"] + rendered = ( + f"{messages[0]['content']}" + f"{function['name']}" + f"{function['arguments']}" + ) + if not tokenize: + return rendered + return [1, 10, 20, 25, 30, 31, 26] + assert add_generation_prompt + return [1, 10] + + tokenized = history.tokenize(tokenizer=Tokenizer()) + + assert tokenized.token_ids == [1, 10, 20, 25, 30, 31, 26] + assert tokenized.flags[2:6] == [art.TokenFlag.SAMPLED] * 4 + assert all(math.isnan(value) for value in tokenized.logprobs[2:6]) + assert not tokenized.flags[0] & art.TokenFlag.SAMPLED + + +@pytest.mark.parametrize( + ("name", "arguments"), + [ + ("lookup", "{}"), + ("art_trajectory_probe_0", '{"art_trajectory_probe":true}'), + ], +) +def test_tool_call_probe_handles_contextual_tokenization( + name: str, arguments: str ) -> None: exchange = _chat_exchange([], []) - logprobs = exchange.response.choices[0].logprobs - assert logprobs is not None - exchange.response.choices[0].logprobs = logprobs.model_copy( - update={ - "content": [ - ChatCompletionTokenLogprob( - token="answer", - logprob=-0.7, - bytes=list(b"answer"), - top_logprobs=[], - ) - ] - } - ) + data = exchange.response.model_dump(mode="python") + choice = data["choices"][0] + choice.pop("prompt_token_ids") + choice.pop("token_ids") + choice["logprobs"] = None + choice["message"] = { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call-1", + "type": "function", + "function": {"name": name, "arguments": arguments}, + } + ], + } + exchange.response = ChatCompletion.model_validate(data) + history = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[exchange]) + ).chat_completions_history() class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> list[int]: + del kwargs + return { + "turn 0": [1], + name: [90], + arguments: [91], + }[text] + def apply_chat_template( self, messages: list[dict[str, Any]], **kwargs: object ) -> list[int]: del kwargs - return [10, 11, 12] if messages[-1]["role"] == "assistant" else [10] + if messages[-1]["role"] != "assistant": + return [1, 10] + function = messages[-1]["tool_calls"][0]["function"] + if function["name"] != name: + return [1, 10, 98, 25, 99, 26] + return [1, 10, 20, 25, 30, 26] - def __call__(self, text: str, **kwargs: object) -> SimpleNamespace: - del text, kwargs - return SimpleNamespace(input_ids=[11]) + tokenized = history.tokenize(tokenizer=Tokenizer()) - monkeypatch.setattr( - "art.trajectories._tokenize._load_tokenizer", lambda _config: Tokenizer() - ) - result = art.tokenize_trajectory( - art.Trajectory(exchanges=TrajectoryExchanges(chat_completions=[exchange])), - base_model="base/model", - ) - assert result.token_ids == [10, 11, 12] - assert result.logprobs[1] == -0.7 - assert math.isnan(result.logprobs[2]) + assert tokenized.token_ids == [1, 10, 20, 25, 30, 26] + assert tokenized.flags[2:5] == [art.TokenFlag.SAMPLED] * 3 + assert all(math.isnan(value) for value in tokenized.logprobs[2:5]) -def test_ambiguous_visible_logprobs_fail_closed( - monkeypatch: pytest.MonkeyPatch, -) -> None: +def test_rerender_proves_each_reasoning_and_content_part_separately() -> None: exchange = _chat_exchange([], []) - logprobs = exchange.response.choices[0].logprobs - assert logprobs is not None - exchange.response.choices[0].logprobs = logprobs.model_copy( - update={ - "content": [ - ChatCompletionTokenLogprob( - token="answer", - logprob=-0.7, - bytes=list(b"answer"), - top_logprobs=[], - ) - ] - } - ) + data = exchange.response.model_dump(mode="python") + choice = data["choices"][0] + choice.pop("prompt_token_ids") + choice.pop("token_ids") + choice["message"] = { + "role": "assistant", + "reasoning": "think", + "content": "answer", + } + choice["logprobs"] = { + "content": [ + { + "token": "answer", + "logprob": -0.7, + "bytes": list(b"answer"), + "top_logprobs": [], + } + ], + "refusal": None, + } + exchange.response = ChatCompletion.model_validate(data) + history = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[exchange]) + ).chat_completions_history() + history.chat_template = "rerender" + + class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> object: + if kwargs.get("return_offsets_mapping"): + think_start = text.index("think") + think_end = think_start + len("think") + answer_start = text.index("answer") + answer_end = answer_start + len("answer") + return { + "input_ids": [1, 10, 11, 12, 13, 14], + "offset_mapping": [ + (0, text.index("")), + (text.index(""), think_start), + (think_start, think_end), + (think_end, answer_start), + (answer_start, answer_end), + (answer_end, len(text)), + ], + } + return {"turn 0": [1], "think": [11], "answer": [13]}[text] + + def apply_chat_template( + self, messages: list[dict[str, Any]], **kwargs: object + ) -> object: + tokenize = kwargs.pop("tokenize") + del kwargs + assistant = messages[-1] + reasoning = assistant.get("reasoning") or assistant.get("reasoning_content") + rendered = ( + f"{messages[0]['content']}" + + (f"{reasoning}" if reasoning else "") + + f"{assistant['content']}" + ) + if not tokenize: + return rendered + if not reasoning: + return [1, 10, 13, 14] + return [1, 10, 11, 12, 13, 14] + + tokenized = history.tokenize(tokenizer=Tokenizer()) + + assert tokenized.flags == [ + art.TokenFlag(0), + art.TokenFlag(0), + art.TokenFlag.SAMPLED, + art.TokenFlag(0), + art.TokenFlag.SAMPLED, + art.TokenFlag(0), + ] + assert tokenized.logprobs[4] == -0.7 + + +def test_rerender_rejects_unproved_multi_part_boundaries() -> None: + exchange = _chat_exchange([], []) + data = exchange.response.model_dump(mode="python") + data["choices"][0].pop("prompt_token_ids") + data["choices"][0].pop("token_ids") + data["choices"][0]["logprobs"] = None + data["choices"][0]["message"] = { + "role": "assistant", + "reasoning": "think", + "content": "answer", + } + exchange.response = ChatCompletion.model_validate(data) + history = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[exchange]) + ).chat_completions_history() + history.chat_template = "rerender" + + class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> list[int]: + del kwargs + return {"turn 0": [1], "think": [11], "answer": [12]}[text] + + def apply_chat_template( + self, messages: list[dict[str, Any]], **kwargs: object + ) -> list[int]: + del messages, kwargs + return [1, 10, 11, 12, 13, 14] + + with pytest.raises(ValueError, match="sampled content boundary"): + history.tokenize(tokenizer=Tokenizer()) + + +def test_empty_sampled_messages_need_no_content_boundary() -> None: + exchange = _chat_exchange([], []) + data = exchange.response.model_dump(mode="python") + data["choices"][0].pop("prompt_token_ids") + data["choices"][0].pop("token_ids") + data["choices"][0]["logprobs"] = None + data["choices"][0]["message"]["content"] = "" + exchange.response = ChatCompletion.model_validate(data) + history = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[exchange]) + ).chat_completions_history() class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> list[int]: + del text, kwargs + return [1] + def apply_chat_template( self, messages: list[dict[str, Any]], **kwargs: object ) -> list[int]: - del kwargs - return [10, 11, 12, 11] if messages[-1]["role"] == "assistant" else [10] + del messages, kwargs + return [1, 2] - def __call__(self, text: str, **kwargs: object) -> SimpleNamespace: - del text, kwargs - return SimpleNamespace(input_ids=[11]) + tokenized = history.tokenize(tokenizer=Tokenizer()) - monkeypatch.setattr( - "art.trajectories._tokenize._load_tokenizer", lambda _config: Tokenizer() - ) - result = art.tokenize_trajectory( - art.Trajectory(exchanges=TrajectoryExchanges(chat_completions=[exchange])), - base_model="base/model", - ) + assert tokenized.token_ids == [1, 2] + assert tokenized.flags == [art.TokenFlag(0), art.TokenFlag(0)] - assert result.token_ids == [10, 11, 12, 11] - assert all(math.isnan(logprob) for logprob in result.logprobs[1:]) +def test_empty_sampled_message_inserts_exact_control_token() -> None: + exchange = _chat_exchange([], [2]) + exchange.response.choices[0].message.content = "" + history = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[exchange]) + ).chat_completions_history() -def test_legacy_logprob_mismatch_fails_closed() -> None: - exchange = _chat_exchange([1], [2, 3]) - choice = exchange.response.choices[0] - assert choice.logprobs is not None - content = choice.logprobs.content - assert content - choice.logprobs = choice.logprobs.model_copy( - update={ - "content": [ - content[0].model_copy( - update={"token": "answer", "bytes": list(b"answer")} - ) - ] - } - ) + class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> list[int]: + del text, kwargs + return [] - result = art.tokenize_trajectory( - art.Trajectory(messages_and_choices=[choice]), - ) + def apply_chat_template( + self, + messages: list[dict[str, Any]], + *, + add_generation_prompt: bool, + **kwargs: object, + ) -> list[int]: + del kwargs + if messages[-1]["role"] == "assistant": + return [1, 99, 9] + assert add_generation_prompt + return [1, 99] - assert result.token_ids == [1, 2, 3] - assert len(result.logprobs) == len(result.token_ids) - assert all(math.isnan(logprob) for logprob in result.logprobs) + tokenized = history.tokenize(tokenizer=Tokenizer()) + assert tokenized.token_ids == [1, 99, 2, 9] + assert tokenized.logprobs[2] == -0.2 + assert tokenized.flags[2] == art.TokenFlag.EXACT | art.TokenFlag.SAMPLED -def test_anthropic_fallback_rejects_unknown_content_blocks( - monkeypatch: pytest.MonkeyPatch, -) -> None: - response = Message.model_validate( - { - "id": "msg_1", - "type": "message", - "role": "assistant", - "model": "test/model", - "content": [{"type": "text", "text": "answer"}], - "stop_reason": "end_turn", - "stop_sequence": None, - "usage": {"input_tokens": 1, "output_tokens": 1}, - } - ) - start = datetime(2026, 1, 1) - image: ImageBlockParam = { - "type": "image", - "source": { - "type": "base64", - "media_type": "image/png", - "data": "...", - }, + +def test_renderer_ignored_refusal_is_appended_for_tokenization() -> None: + exchange = _chat_exchange([], []) + data = exchange.response.model_dump(mode="python") + choice = data["choices"][0] + choice.pop("prompt_token_ids") + choice.pop("token_ids") + choice["message"] = { + "role": "assistant", + "content": "answer", + "refusal": "declined", } - message: MessageParam = {"role": "user", "content": [image]} - exchange = MessagesExchange( - request=MessagesRequest( - model="test/model", - messages=[message], - ), - response=response, - start_time=start, - end_time=start + timedelta(seconds=1), - ) - monkeypatch.setattr( - "art.trajectories._tokenize._load_tokenizer", lambda _config: _FakeTokenizer() - ) + choice["logprobs"] = { + "content": [ + { + "token": "answer", + "logprob": -0.4, + "bytes": list(b"answer"), + "top_logprobs": [], + } + ], + "refusal": [ + { + "token": "declined", + "logprob": -0.5, + "bytes": list(b"declined"), + "top_logprobs": [], + } + ], + } + exchange.response = ChatCompletion.model_validate(data) + history = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[exchange]) + ).chat_completions_history() - with pytest.raises(ValueError, match="Unsupported Anthropic content block"): - art.tokenize_trajectory( - art.Trajectory(exchanges=TrajectoryExchanges(messages=[exchange])), - base_model="base/model", - ) + class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> object: + if kwargs.get("return_offsets_mapping"): + answer_start = text.index("answer") + declined_start = text.index("declined") + declined_end = declined_start + len("declined") + return { + "input_ids": [1, 4, 5], + "offset_mapping": [ + (0, answer_start), + (answer_start, declined_start), + (declined_start, declined_end), + ], + } + return { + "turn 0": [1], + "answer": [4], + "declined": [5], + "answerdeclined": [4, 5], + }[text] + + def apply_chat_template( + self, messages: list[dict[str, Any]], **kwargs: object + ) -> object: + tokenize = kwargs.pop("tokenize") + del kwargs + if not tokenize: + return "".join( + f"{message.get('content') or ''}" + for message in messages + ) + result: list[int] = [] + for message in messages: + encoded = self(str(message.get("content") or "")) + assert isinstance(encoded, list) + for token in encoded: + assert isinstance(token, int) + result.append(token) + return result + tokenized = history.tokenize(tokenizer=Tokenizer()) -def test_undecodable_visible_token_bytes_fall_back_to_nan( - monkeypatch: pytest.MonkeyPatch, -) -> None: + assert tokenized.token_ids == [1, 4, 5] + assert tokenized.logprobs[1:] == [-0.4, -0.5] + assert tokenized.flags[1:] == [art.TokenFlag.SAMPLED] * 2 + + +def test_renderer_reasoning_content_alias_preserves_reasoning() -> None: exchange = _chat_exchange([], []) - logprobs = exchange.response.choices[0].logprobs - assert logprobs is not None - exchange.response.choices[0].logprobs = logprobs.model_copy( - update={ - "content": [ - ChatCompletionTokenLogprob( - token="ordinary-token", - logprob=-0.7, - bytes=[0xF0], - top_logprobs=[], - ) - ] - } - ) + data = exchange.response.model_dump(mode="python") + choice = data["choices"][0] + choice.pop("prompt_token_ids") + choice.pop("token_ids") + choice["logprobs"] = None + choice["message"] = { + "role": "assistant", + "reasoning": "think", + "content": "answer", + } + exchange.response = ChatCompletion.model_validate(data) + history = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[exchange]) + ).chat_completions_history() class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> list[int]: + del kwargs + return {"turn 0": [1], "think": [2], "answer": [3]}[text] + def apply_chat_template( self, messages: list[dict[str, Any]], **kwargs: object ) -> list[int]: del kwargs - return [10, 11] if messages[-1]["role"] == "assistant" else [10] + result: list[int] = [] + for message in messages: + if reasoning := message.get("reasoning_content"): + result.extend(self(str(reasoning))) + if content := message.get("content"): + result.extend(self(str(content))) + return result - monkeypatch.setattr( - "art.trajectories._tokenize._load_tokenizer", lambda _config: Tokenizer() - ) - result = art.tokenize_trajectory( - art.Trajectory(exchanges=TrajectoryExchanges(chat_completions=[exchange])), - base_model="base/model", - ) + tokenized = history.tokenize(tokenizer=Tokenizer()) - assert result.token_ids == [10, 11] - assert math.isnan(result.logprobs[1]) + assert tokenized.token_ids == [1, 2, 3] + assert tokenized.flags == [ + art.TokenFlag(0), + art.TokenFlag.SAMPLED, + art.TokenFlag.SAMPLED, + ] -def test_json_round_trip_preserves_exchange_types() -> None: - exchange = _chat_exchange([1], [2]) - request: dict[str, Any] = { - "model": "test/model", - "messages": [ - {"role": "assistant", "content": "answer", "reasoning": "thinking"} +def test_trimmed_render_preserves_authoritative_textual_logprob_tokens() -> None: + exchange = _chat_exchange([], []) + data = exchange.response.model_dump(mode="python") + choice = data["choices"][0] + choice.pop("prompt_token_ids") + choice.pop("token_ids") + choice["message"]["content"] = " helloworld " + choice["logprobs"] = { + "content": [ + { + "token": " hello", + "logprob": -0.4, + "bytes": list(b" hello"), + "top_logprobs": [], + }, + { + "token": "world", + "logprob": -0.45, + "bytes": list(b"world"), + "top_logprobs": [], + }, + { + "token": " ", + "logprob": -0.5, + "bytes": [32], + "top_logprobs": [], + }, ], + "refusal": None, } - exchange.request = ChatCompletionsRequest(**request) - original = art.Trajectory( + exchange.response = ChatCompletion.model_validate(data) + history = art.Trajectory( exchanges=TrajectoryExchanges(chat_completions=[exchange]) - ) - dumped = original.model_dump(mode="json", warnings="error") - assert dumped["exchanges"]["chat_completions"][0]["request"] == request - restored = art.Trajectory.model_validate_json(original.model_dump_json()) - assert restored.model_dump(mode="json") == original.model_dump(mode="json") - assert isinstance(restored.exchanges.chat_completions[0].response, ChatCompletion) - + ).chat_completions_history() -def _response_exchange( - response_id: str, - output_id: int, - *, - previous_response_id: str | None = None, - offset: int = 0, -) -> ResponsesExchange: - response = Response.model_validate( - { - "id": response_id, - "created_at": float(offset), - "model": "test/model", - "object": "response", - "output": [ - { - "id": f"message-{response_id}", - "type": "message", - "role": "assistant", - "status": "completed", - "content": [ - { - "type": "output_text", - "text": "answer", - "annotations": [], - "logprobs": [], - } + class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> object: + if kwargs.get("return_offsets_mapping"): + if text == " helloworld ": + return { + "input_ids": [1118, 2222, 220], + "offset_mapping": [(0, 6), (6, 11), (11, 12)], + } + content_start = text.index("helloworld") + content_end = content_start + len("helloworld") + return { + "input_ids": [1, 3765], + "offset_mapping": [ + (0, content_start), + (content_start, content_end), ], } - ], - "parallel_tool_calls": True, - "tool_choice": "auto", - "tools": [], - "raw_output_tokens": [{"token_id": output_id, "logprob": -0.1}], - } - ) - request = ResponsesRequest(model="test/model", input=f"turn {offset}") - if previous_response_id is not None: - request["previous_response_id"] = previous_response_id - start = datetime(2026, 1, 1) + timedelta(seconds=offset) - return ResponsesExchange( - request=request, - response=response, - start_time=start, - end_time=start + timedelta(milliseconds=1), - ) + return { + "turn 0": [1], + " helloworld ": [1118, 2222, 220], + " hello": [1118], + "world": [3333], + " ": [220], + }[text] + def apply_chat_template( + self, messages: list[dict[str, Any]], **kwargs: object + ) -> object: + tokenize = kwargs.pop("tokenize") + del kwargs + rendered = ( + f"{messages[0]['content']}" + f"{str(messages[-1]['content']).strip()}" + ) + if not tokenize: + return rendered + return [1, 3765] -def _response_with_content_logprobs(*, exact_second: bool) -> ResponsesExchange: - exchange = _response_exchange("response-content-logprobs", 0) - data = exchange.response.model_dump(mode="python") - data.pop("raw_output_tokens", None) + tokenized = history.tokenize(tokenizer=Tokenizer()) - def entry(token: str, token_id: int | None, logprob: float) -> dict[str, Any]: - return { - "token": token, - "logprob": logprob, - "bytes": list(("a" if token_id == 11 else "b").encode()), + assert tokenized.token_ids == [1, 1118, 2222, 220] + assert tokenized.logprobs[1:] == [-0.4, -0.45, -0.5] + assert tokenized.flags[1:] == [art.TokenFlag.SAMPLED] * 3 + + +def test_textual_logprobs_reconstruct_split_utf8_bytes() -> None: + from art.trajectories._tokenize import _visible_token_evidence + + exchange = _chat_exchange([], []) + data = exchange.response.model_dump(mode="python") + choice = data["choices"][0] + choice.pop("prompt_token_ids") + choice.pop("token_ids") + choice["message"]["content"] = "😊" + choice["logprobs"]["content"] = [ + { + "token": "�", + "logprob": -0.4, + "bytes": [240, 159, 152], "top_logprobs": [], - **({"token_id": token_id} if token_id is not None else {}), - } + }, + { + "token": "�", + "logprob": -0.5, + "bytes": [138], + "top_logprobs": [], + }, + ] + exchange.response = ChatCompletion.model_validate(data) - data["output"][0]["content"] = [ + class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> list[int]: + del kwargs + assert text == "😊" + return [11, 12] + + def apply_chat_template( + self, + messages: list[dict[str, Any]], + *, + tools: object, + tokenize: bool, + add_generation_prompt: bool, + chat_template: str | None = None, + **kwargs: object, + ) -> object: + del messages, tools, tokenize, add_generation_prompt, chat_template, kwargs + raise AssertionError("template rendering is not expected") + + assert _visible_token_evidence(Tokenizer(), exchange, sampled_text="😊") == ( + [11, 12], + [-0.4, -0.5], + ) + + +def test_textual_logprobs_reject_shifted_contextual_token_boundaries() -> None: + from art.trajectories._tokenize import _visible_token_evidence + + exchange = _chat_exchange([], []) + data = exchange.response.model_dump(mode="python") + choice = data["choices"][0] + choice.pop("prompt_token_ids") + choice.pop("token_ids") + choice["message"]["content"] = " penalates" + choice["logprobs"]["content"] = [ { - "type": "output_text", - "text": "a", - "annotations": [], - "logprobs": [entry("token_id:11", 11, -0.1)], + "token": " pena", + "logprob": -0.4, + "bytes": list(b" pena"), + "top_logprobs": [], }, { - "type": "output_text", - "text": "b", - "annotations": [], - "logprobs": [ - entry( - "token_id:12" if exact_second else "b", - 12 if exact_second else None, - -0.2, - ) - ], + "token": "lates", + "logprob": -0.5, + "bytes": list(b"lates"), + "top_logprobs": [], }, ] - exchange.response = Response.model_validate(data) - return exchange + exchange.response = ChatCompletion.model_validate(data) + class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> object: + if kwargs.get("return_offsets_mapping"): + return { + "input_ids": [30, 31], + "offset_mapping": [(0, 6), (6, 10)], + } + return {" pena": [11], "lates": [12]}[text] -def test_responses_aggregates_complete_exact_pairs_across_content_blocks( - monkeypatch: pytest.MonkeyPatch, + def apply_chat_template( + self, + messages: list[dict[str, Any]], + *, + tools: object, + tokenize: bool, + add_generation_prompt: bool, + chat_template: str | None = None, + **kwargs: object, + ) -> object: + del messages, tools, tokenize, add_generation_prompt, chat_template, kwargs + raise AssertionError("template rendering is not expected") + + assert _visible_token_evidence( + Tokenizer(), exchange, sampled_text=" penalates" + ) == ([11, 12], [-0.4, -0.5]) + + +@pytest.mark.parametrize( + ("content", "token_id", "trim", "adjacent_scaffold", "reject_empty"), + [ + (" ", 220, True, False, False), + ("\n", 198, True, False, False), + (" ", 220, False, False, False), + (" ", 220, False, True, False), + (" ", 220, False, False, True), + ], +) +def test_trimmed_whitespace_output_inserts_authoritative_logprob_token( + content: str, + token_id: int, + trim: bool, + adjacent_scaffold: bool, + reject_empty: bool, ) -> None: + exchange = _chat_exchange([], []) + data = exchange.response.model_dump(mode="python") + choice = data["choices"][0] + choice.pop("prompt_token_ids") + choice.pop("token_ids") + choice["message"]["content"] = content + choice["logprobs"] = { + "content": [ + { + "token": content, + "logprob": -0.5, + "bytes": list(content.encode()), + "top_logprobs": [], + } + ], + "refusal": None, + } + exchange.response = ChatCompletion.model_validate(data) + history = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[exchange]) + ).chat_completions_history() + class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> object: + if kwargs.get("return_offsets_mapping"): + start = text.index("") + len("") + end = text.index("") + if start != end: + scaffold_start = end - int(adjacent_scaffold) + return { + "input_ids": [ + 1, + token_id, + *([token_id] if adjacent_scaffold else []), + 9, + ], + "offset_mapping": [ + (0, start), + (start, scaffold_start), + *([(scaffold_start, end)] if adjacent_scaffold else []), + (end, len(text)), + ], + } + return { + "input_ids": [1, 9], + "offset_mapping": [(0, start), (start, len(text))], + } + return {"turn 0": [1], content: [token_id]}[text] + def apply_chat_template( self, messages: list[dict[str, Any]], **kwargs: object - ) -> list[int]: - del messages, kwargs - return [10] + ) -> object: + tokenize = kwargs.pop("tokenize") + del kwargs + content_value = str(messages[-1]["content"]) + if trim: + content_value = content_value.strip() + scaffold = " " if adjacent_scaffold else "" + rendered = ( + f"{messages[0]['content']}" + f"{content_value}{scaffold}" + ) + if not tokenize: + return rendered + if reject_empty and not content_value: + raise RuntimeError("template rejects empty assistant content") + rendered_id = 777 if "ART_TRAJECTORY" in content_value else token_id + return [ + 1, + *([rendered_id] if content_value else []), + *([token_id] if adjacent_scaffold else []), + 9, + ] - monkeypatch.setattr( - "art.trajectories._tokenize._load_tokenizer", lambda _config: Tokenizer() - ) - result = art.tokenize_trajectory( - art.Trajectory( - exchanges=TrajectoryExchanges( - responses=[_response_with_content_logprobs(exact_second=True)] + tokenized = history.tokenize(tokenizer=Tokenizer()) + + assert tokenized.token_ids == [ + 1, + token_id, + *([token_id] if adjacent_scaffold else []), + 9, + ] + assert tokenized.logprobs[1] == -0.5 + assert tokenized.flags[1] == art.TokenFlag.SAMPLED + if adjacent_scaffold: + assert tokenized.flags[2] == art.TokenFlag(0) + + +def _repeated_text_rerender_history(turn_count: int) -> art.ChatCompletionsHistory: + exchanges: list[ChatCompletionsExchange] = [] + prompt: list[int] = [] + messages: list[ChatCompletionMessageParam] = [] + for index in range(turn_count): + prompt.extend([index * 2]) + messages.append({"role": "user", "content": f"u{index}"}) + exchange = _chat_exchange(list(prompt), [index * 2 + 1], offset=index) + exchange.request["messages"] = list(messages) + exchanges.append(exchange) + prompt.append(index * 2 + 1) + messages.append({"role": "assistant", "content": "answer"}) + history = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=exchanges) + ).chat_completions_history() + history.chat_template = "rerender" + return history + + +class _RepeatedTextTokenizer: + def __init__(self) -> None: + self.apply_calls = 0 + + def __call__(self, text: str, **kwargs: object) -> object: + if not kwargs.get("return_offsets_mapping"): + return [1000] if text == "answer" else [2000 + int(text[1:])] + token_ids: list[int] = [] + offsets: list[tuple[int, int]] = [] + for match in re.finditer(r"<([ua])>(.*?)", text): + content_start, content_end = match.span(2) + token_ids.extend( + [ + 3000 if match.group(1) == "u" else 3001, + 1000 + if match.group(2) == "answer" + else 2000 + int(match.group(2)[1:]), + 3002, + ] ) - ), - base_model="base/model", - ) + offsets.extend( + [ + (match.start(), content_start), + (content_start, content_end), + (content_end, match.end()), + ] + ) + return {"input_ids": token_ids, "offset_mapping": offsets} + + def apply_chat_template( + self, + messages: list[dict[str, Any]], + *, + tokenize: bool, + **kwargs: object, + ) -> object: + del kwargs + self.apply_calls += 1 + rendered = "".join( + f"<{'a' if message['role'] == 'assistant' else 'u'}>" + f"{message['content']}" + for message in messages + ) + return self(rendered, return_offsets_mapping=True) if tokenize else rendered - assert result.token_ids == [10, 11, 12] - assert result.logprobs[1:] == [-0.1, -0.2] +def test_rerender_calls_chat_template_once_for_many_turns() -> None: + history = _repeated_text_rerender_history(32) -def test_responses_empty_raw_tokens_fall_back_for_visible_output( + tokenizer = _RepeatedTextTokenizer() + history.tokenize(tokenizer=tokenizer) + + assert tokenizer.apply_calls == 2 + + +def test_repeated_text_rerender_scaling_is_near_linear() -> None: + medians: list[float] = [] + for turn_count in (32, 64, 128): + history = _repeated_text_rerender_history(turn_count) + samples: list[float] = [] + for _ in range(5): + tokenizer = _RepeatedTextTokenizer() + started = perf_counter() + history.tokenize(tokenizer=tokenizer) + samples.append(perf_counter() - started) + assert tokenizer.apply_calls == 2 + medians.append(median(samples)) + + assert medians[1] < medians[0] * 3 + assert medians[2] < medians[1] * 3 + + +def test_reasoning_split_trajectory_reuses_prevalidated_projections( monkeypatch: pytest.MonkeyPatch, ) -> None: - exchange = _response_exchange("response-empty-raw", 0) - data = exchange.response.model_dump(mode="python") - data["raw_output_tokens"] = [] - exchange.response = Response.model_validate(data) - monkeypatch.setattr( - "art.trajectories._tokenize._load_tokenizer", lambda _config: _FakeTokenizer() + exchanges: list[ChatCompletionsExchange] = [] + request_messages: list[ChatCompletionMessageParam] = [] + prompt: list[int] = [] + for index in range(40): + request_messages.append({"role": "user", "content": f"u{index}"}) + prompt.append(3000 + index) + exchange = _chat_exchange( + list(prompt), [1000 + index, 2000 + index], offset=index + ) + exchange.request["messages"] = list(request_messages) + payload = exchange.response.model_dump(mode="python") + payload["choices"][0]["message"] = { + "role": "assistant", + "reasoning": f"r{index}", + "content": f"a{index}", + } + exchange.response = ChatCompletion.model_validate(payload) + exchanges.append(exchange) + request_messages.append({"role": "assistant", "content": f"a{index}"}) + prompt.append(2000 + index) + trajectory = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=exchanges) ) - - result = art.tokenize_trajectory( - art.Trajectory(exchanges=TrajectoryExchanges(responses=[exchange])), - base_model="base/model", - chat_template="template", - chat_template_kwargs={}, + projected = trajectory.chat_completions_histories() + assert len(projected) == 40 + for history in projected: + for source in history.message_sources: + if source is not None: + assert any(source.exchange is exchange for exchange in exchanges) + assert all( + any( + source is not None + and source.choice_index == 0 + and source.exchange is exchanges[0] + for source in history.message_sources + ) + for history in projected ) - assert result.token_ids == [10, 11] + from art.trajectories import _tokenize + original = _tokenize._history_matches_projection + calls = 0 -def test_responses_does_not_use_partial_exact_content_pairs( - monkeypatch: pytest.MonkeyPatch, -) -> None: - class Tokenizer: - def apply_chat_template( - self, messages: list[dict[str, Any]], **kwargs: object - ) -> list[int]: - del kwargs - return [10, 11, 12] if messages[-1]["role"] == "assistant" else [10] + def counted(history: art.History) -> bool: + nonlocal calls + calls += 1 + return original(history) - def __call__(self, text: str, **kwargs: object) -> SimpleNamespace: - del kwargs - return SimpleNamespace( - input_ids=[11 if text in {"a", "token_id:11"} else 12] - ) + monkeypatch.setattr(_tokenize, "_history_matches_projection", counted) - monkeypatch.setattr( - "art.trajectories._tokenize._load_tokenizer", lambda _config: Tokenizer() - ) - result = art.tokenize_trajectory( - art.Trajectory( - exchanges=TrajectoryExchanges( - responses=[_response_with_content_logprobs(exact_second=False)] - ) - ), - base_model="base/model", - ) + tokenized = trajectory.tokenize(multi_history=True) + _, traces = _tokenize._tokenize_trajectory_with_trace(trajectory) + first_key = next(key for key in traces[0].source_keys if key is not None) - assert result.token_ids == [10, 11, 12] - assert result.logprobs[1:] == [-0.1, -0.2] + assert len(tokenized.histories) == 40 + assert all(first_key in trace.sources for trace in traces) + assert calls == 0 -def test_responses_rejects_only_unrenderable_prompt_history( - monkeypatch: pytest.MonkeyPatch, -) -> None: +def test_explicit_template_override_rerenders_exact_exchange_scaffold() -> None: + trajectory = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[_chat_exchange([1], [2])]) + ) + class Tokenizer: + def __call__(self, text: str, *, add_special_tokens: bool = False) -> list[int]: + assert not add_special_tokens + return {"turn 0": [10], "answer": [20]}[text] + def apply_chat_template( - self, messages: list[dict[str, Any]], **kwargs: object + self, + messages: list[dict[str, Any]], + *, + add_generation_prompt: bool, + chat_template: str | None = None, + **kwargs: object, ) -> list[int]: del kwargs - assistant_count = sum( - message["role"] == "assistant" for message in messages - ) - return [10, *range(2, 2 + assistant_count)] + assert chat_template == "custom" + if messages[-1]["role"] == "assistant": + if str(messages[-1]["content"]).startswith("ART_TRAJECTORY_"): + return [10, 999, 30] + return [10, 20, 30] + assert add_generation_prompt + return [10] - monkeypatch.setattr( - "art.trajectories._tokenize._load_tokenizer", lambda _config: Tokenizer() + tokenized = trajectory.tokenize( + tokenizer=Tokenizer(), + chat_template="custom", ) - request_reasoning = _response_exchange("request-reasoning", 2) - request_reasoning.request["input"] = [ - { - "id": "reasoning-1", - "summary": [{"type": "summary_text", "text": "request thought"}], - "type": "reasoning", - } - ] - - response_reasoning = _response_exchange("response-reasoning", 2) - data = response_reasoning.response.model_dump(mode="python") - data["output"] = [ - { - "id": "reasoning-2", - "summary": [{"type": "summary_text", "text": "response thought"}], - "type": "reasoning", - } - ] - data.pop("raw_output_tokens", None) - response_reasoning.response = Response.model_validate(data) - art.tokenize_trajectory( - art.Trajectory(exchanges=TrajectoryExchanges(responses=[request_reasoning])), - base_model="base/model", - ) + assert tokenized.token_ids == [10, 2, 30] + assert tokenized.logprobs[1] == -0.2 - single = art.Trajectory( - exchanges=TrajectoryExchanges(responses=[response_reasoning]) - ) - assert art.tokenize_trajectory(single, base_model="base/model").token_ids == [ - 10, - 2, - ] - continuation = _response_exchange( - "continuation", - 3, - previous_response_id=response_reasoning.response.id, - offset=1, +def test_responses_external_context_requires_or_uses_exact_prompt_tokens() -> None: + exchange = _response_exchange( + "external", 2, previous_response_id="outside-trajectory" ) - assert art.tokenize_trajectory( - art.Trajectory( - exchanges=TrajectoryExchanges(responses=[response_reasoning, continuation]) - ), - base_model="base/model", - ).token_ids == [10, 2, 3] - + response = exchange.response.model_dump(mode="python") + response.pop("token_generations", None) + exchange.response = Response.model_validate(response) + history = art.Trajectory( + exchanges=TrajectoryExchanges(responses=[exchange]) + ).responses_history() + with pytest.raises(ValueError, match="without exact prompt tokens"): + history.tokenize(base_model="base/model") -def test_responses_opaque_reasoning_requires_exact_tokens( - monkeypatch: pytest.MonkeyPatch, -) -> None: - exchange = _response_exchange("opaque-reasoning", 2) response = exchange.response.model_dump(mode="python") - response["output"] = [ + response["token_generations"] = [ { - "id": "reasoning-1", - "encrypted_content": "opaque", - "summary": [], - "type": "reasoning", + "prompt_token_ids": [7, 8], + "output_tokens": [{"token_id": 2, "logprob": -0.1}], + "output_indices": [0], } ] exchange.response = Response.model_validate(response) - monkeypatch.setattr( - "art.trajectories._tokenize._load_tokenizer", lambda _config: _FakeTokenizer() + tokenized = ( + art.Trajectory(exchanges=TrajectoryExchanges(responses=[exchange])) + .responses_history() + .tokenize() ) - assert art.tokenize_trajectory( - art.Trajectory(exchanges=TrajectoryExchanges(responses=[exchange])), - base_model="base/model", - ).token_ids == [10, 2] + assert tokenized.token_ids == [7, 8, 2] + assert tokenized.flags == [ + art.TokenFlag.EXACT, + art.TokenFlag.EXACT, + art.TokenFlag.EXACT | art.TokenFlag.SAMPLED, + ] + +def test_responses_conversation_requires_exact_prompt_tokens() -> None: + exchange = _response_exchange("conversation", 2) + exchange.request["conversation"] = "conversation-1" response = exchange.response.model_dump(mode="python") - response.pop("raw_output_tokens", None) + response.pop("token_generations", None) exchange.response = Response.model_validate(response) - with pytest.raises(ValueError, match="no renderable text"): - art.tokenize_trajectory( - art.Trajectory(exchanges=TrajectoryExchanges(responses=[exchange])), - base_model="base/model", - ) + trajectory = art.Trajectory(exchanges=TrajectoryExchanges(responses=[exchange])) + with pytest.raises(ValueError, match="conversation history requires exact"): + trajectory.tokenize(base_model="base/model") -def test_responses_parallel_function_calls_form_one_assistant_turn( - monkeypatch: pytest.MonkeyPatch, -) -> None: - exchange = _response_exchange("parallel-tools", 2) - exchange.request["input"] = [ - { - "id": "reasoning-1", - "summary": [{"type": "summary_text", "text": "think"}], - "type": "reasoning", - }, - {"type": "function_call", "call_id": "one", "name": "first", "arguments": "{}"}, + response = exchange.response.model_dump(mode="python") + response["token_generations"] = [ { - "type": "function_call", - "call_id": "two", - "name": "second", - "arguments": "{}", - }, + "prompt_token_ids": [5], + "output_tokens": [{"token_id": 2, "logprob": -0.1}], + "output_indices": [0], + } ] - seen: list[list[dict[str, Any]]] = [] + exchange.response = Response.model_validate(response) + assert trajectory.tokenize().token_ids == [5, 2] - class Tokenizer(_FakeTokenizer): - def apply_chat_template( - self, messages: list[dict[str, Any]], **kwargs: object - ) -> list[int]: - seen.append(messages) - return super().apply_chat_template(messages, **kwargs) - monkeypatch.setattr( - "art.trajectories._tokenize._load_tokenizer", lambda _config: Tokenizer() +def test_tokenized_results_materialize_metadata_and_group_shape() -> None: + trajectory = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[_chat_exchange([1], [2])]), + reward=0.75, + metrics={"correct": True}, + metadata={"source": {"name": "unit"}}, ) - art.tokenize_trajectory( - art.Trajectory(exchanges=TrajectoryExchanges(responses=[exchange])), - base_model="base/model", + group = art.TrajectoryGroup( + [trajectory], metrics={"batch": 1}, metadata={"split": "test"} ) - assistant = seen[0][0] - assert assistant["reasoning"] == "think" - assert [call["function"]["name"] for call in assistant["tool_calls"]] == [ - "first", - "second", + tokenized = group.tokenize() + + assert tokenized.trajectories[0].model == "test/model" + assert tokenized.trajectories[0].reward == 0.75 + assert tokenized.trajectories[0].metadata == {"source": {"name": "unit"}} + assert tokenized.metrics == {"batch": 1} + assert tokenized.metadata == {"split": "test"} + assert "underlying" not in tokenized.model_dump() + + +def test_private_trace_covers_sampled_tokens_for_every_protocol() -> None: + from art.trajectories._tokenize import _tokenize_trajectory_with_trace + + trajectories = [ + art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[_chat_exchange([1], [2])]) + ), + art.Trajectory( + exchanges=TrajectoryExchanges(completions=[_completion_exchange()]) + ), + art.Trajectory( + exchanges=TrajectoryExchanges( + responses=[_response_exchange("response-trace", 2)] + ) + ), + art.Trajectory( + exchanges=TrajectoryExchanges( + messages=[ + _message_exchange( + MessagesRequest( + model="test/model", + messages=[{"role": "user", "content": "question"}], + max_tokens=16, + ), + identifier="message-trace", + prompt_token_ids=[1], + token_ids=[2], + logprobs=[-0.2], + ) + ] + ) + ), ] + for trajectory in trajectories: + tokenized, traces = _tokenize_trajectory_with_trace(trajectory) -def test_tokenization_rejects_mutated_mixed_representation() -> None: + assert len(tokenized.histories) == len(traces) == 1 + history = tokenized.histories[0] + trace = traces[0] + trace.validate(history) + assert sum(key is not None for key in trace.source_keys) == sum( + bool(flag & art.TokenFlag.SAMPLED) for flag in history.flags + ) + assert len(trace.sources) == 1 + + +def test_private_trace_keys_do_not_collide_for_repeated_empty_response_ids() -> None: + from art.trajectories._tokenize import _tokenize_trajectory_with_trace + + first = _chat_exchange([1], [2]) + second = _chat_exchange([1, 2, 3], [4], offset=1) + first.response.id = second.response.id = "" + second.start_time = first.start_time + second.end_time = first.end_time trajectory = art.Trajectory( - messages_and_choices=[{"role": "user", "content": "hi"}] + exchanges=TrajectoryExchanges(chat_completions=[first, second]) ) - trajectory.exchanges.chat_completions.append(_chat_exchange([1], [2])) - with pytest.raises(ValueError, match="both exchanges and legacy histories"): - art.tokenize_trajectory(trajectory) + tokenized, [trace] = _tokenize_trajectory_with_trace(trajectory) + sampled_keys = [key for key in trace.source_keys if key is not None] + assert tokenized.histories[0].token_ids == [1, 2, 3, 4] + assert len(set(sampled_keys)) == 2 + assert len(trace.sources) == 2 + + +def test_responses_fallback_trace_does_not_retrain_echoed_output_items() -> None: + from art.trajectories._tokenize import ( + _first_introduction_mask, + _tokenize_trajectory_with_trace, + ) + + def output(item_id: str, text: str) -> ResponseOutputMessageParam: + return { + "id": item_id, + "type": "message", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": text, "annotations": []}], + } + + def response( + response_id: str, items: list[ResponseOutputMessageParam], offset: int + ) -> Response: + return Response.model_validate( + { + "id": response_id, + "created_at": float(offset), + "model": "test/model", + "object": "response", + "output": items, + "parallel_tool_calls": True, + "tool_choice": "auto", + "tools": [], + } + ) + + user = EasyInputMessageParam(role="user", content="question") + first_outputs = [output("one", "first"), output("two", "second")] + first_response = response( + "response-fallback", + first_outputs, + 0, + ) + first_input: ResponseInputParam = [user] + first_request = ResponsesRequest(model="test/model", input=first_input) + first_request["chat_template_kwargs"] = {"enable_thinking": False} + first = ResponsesExchange( + request=first_request, + response=first_response, + start_time=datetime(2026, 1, 1), + end_time=datetime(2026, 1, 1, 0, 0, 0, 1000), + ) + echoed: ResponseInputParam = [*first_outputs] + second_input: ResponseInputParam = [ + user, + *echoed, + EasyInputMessageParam(role="user", content="continue"), + ] + second = ResponsesExchange( + request=ResponsesRequest( + model="test/model", + input=second_input, + ), + response=response("response-final", [output("final", "final")], 1), + start_time=datetime(2026, 1, 1, 0, 0, 1), + end_time=datetime(2026, 1, 1, 0, 0, 1, 1000), + ) -def test_responses_previous_response_id_resolves_local_history( - monkeypatch: pytest.MonkeyPatch, -) -> None: class Tokenizer: + def __call__(self, text: str, **kwargs: object) -> list[int]: + del kwargs + return { + "question": [1], + "first": [2], + "second": [3], + "firstsecond": [2, 3], + "continue": [4], + "final": [5], + }[text] + def apply_chat_template( self, messages: list[dict[str, Any]], **kwargs: object ) -> list[int]: del kwargs - return [10] if len(messages) == 1 else [10, 20, 11] + return [ + token for message in messages for token in self(str(message["content"])) + ] - monkeypatch.setattr( - "art.trajectories._tokenize._load_tokenizer", lambda _config: Tokenizer() - ) - first = _response_exchange("resp-1", 20) - second = _response_exchange("resp-2", 30, previous_response_id="resp-1", offset=1) trajectory = art.Trajectory( exchanges=TrajectoryExchanges(responses=[first, second]) ) + tokenized, traces = _tokenize_trajectory_with_trace( + trajectory, + tokenizer=Tokenizer(), + chat_template_kwargs={"enable_thinking": False}, + ) - assert art.tokenize_trajectory(trajectory, base_model="base/model").token_ids == [ - 10, - 20, - 11, - 30, - ] + assert len(tokenized.histories) == len(traces) == 2 + seen: set[object] = set() + trained_first_response = 0 + first_response_indices: set[int] = set() + for trace in traces: + trainable = _first_introduction_mask(trace.source_keys, seen) + for selected, key in zip(trainable, trace.source_keys, strict=True): + if key is not None and key.response_id == first_response.id: + first_response_indices.add(key.index) + trained_first_response += selected + + assert first_response_indices == {0} + assert trained_first_response == 2 - second.request["previous_response_id"] = "missing" - with pytest.raises(ValueError, match="outside this trajectory"): - art.tokenize_trajectory(trajectory, base_model="base/model") + +def test_completions_history_requires_exhaustive_source_spans() -> None: + history = art.CompletionsTokenHistory( + model="test/model", + prompt=[1], + prompt_sources=[], + sampled_spans=[], + ) + + with pytest.raises(ValueError, match="exhaustively cover"): + history.tokenize() def test_exchange_trajectories_feed_existing_training_tokenizer( @@ -1028,6 +5369,10 @@ def apply_chat_template( self.calls.append(kwargs) return [1, 2] if messages[-1]["role"] == "assistant" else [1] + def __call__(self, text: str, *, add_special_tokens: bool = False) -> list[int]: + del text, add_special_tokens + return [2] + def decode(self, token_id: int) -> str: return str(token_id) diff --git a/tests/unit/trajectories/test_tokenized_models.py b/tests/unit/trajectories/test_tokenized_models.py new file mode 100644 index 000000000..97c7eaca0 --- /dev/null +++ b/tests/unit/trajectories/test_tokenized_models.py @@ -0,0 +1,147 @@ +import math + +import art + + +def _history() -> art.TokenizedHistory: + return art.TokenizedHistory( + model="policy", + token_ids=[1, 2], + logprobs=[math.nan, -0.25], + flags=[art.TokenFlag.EXACT, art.TokenFlag.EXACT | art.TokenFlag.SAMPLED], + ) + + +def _assert_history_round_trip( + restored: art.TokenizedHistory, expected: art.TokenizedHistory +) -> None: + assert restored.model == expected.model + assert restored.token_ids == expected.token_ids + assert restored.flags == expected.flags + assert math.isnan(restored.logprobs[0]) + assert restored.logprobs[1:] == expected.logprobs[1:] + + +def test_tokenized_history_nan_json_round_trip() -> None: + value = _history() + payload = value.model_dump_json() + assert '"NaN"' in payload + restored = art.TokenizedHistory.model_validate_json(payload) + _assert_history_round_trip(restored, value) + + +def test_tokenized_trajectory_nan_json_round_trip() -> None: + value = art.TokenizedTrajectory( + **_history().model_dump(), + reward=1.0, + metrics={"count": 1}, + metadata={"source": "test"}, + ) + restored = art.TokenizedTrajectory.model_validate_json(value.model_dump_json()) + _assert_history_round_trip(restored, value) + assert restored.reward == value.reward + assert restored.metrics == value.metrics + assert restored.metadata == value.metadata + + +def test_nested_tokenized_models_nan_json_round_trip() -> None: + trajectory = art.TokenizedMultiHistoryTrajectory( + histories=[_history()], + reward=1.0, + metrics={}, + metadata={}, + ) + group = art.TokenizedTrajectoryGroup[art.TokenizedMultiHistoryTrajectory]( + trajectories=[trajectory], + metrics={}, + metadata={}, + ) + restored = art.TokenizedTrajectoryGroup[ + art.TokenizedMultiHistoryTrajectory + ].model_validate_json(group.model_dump_json()) + restored_trajectory = restored.trajectories[0] + _assert_history_round_trip(restored_trajectory.histories[0], _history()) + assert restored_trajectory.reward == trajectory.reward + assert restored.metrics == group.metrics + assert restored.metadata == group.metadata + + +def test_public_group_tokenization_nan_json_round_trip() -> None: + from datetime import datetime + + from openai.types.chat import ChatCompletion + + from art.trajectories import ( + ChatCompletionsExchange, + ChatCompletionsRequest, + TrajectoryExchanges, + ) + + exchange = ChatCompletionsExchange( + request=ChatCompletionsRequest( + model="policy", + messages=[{"role": "user", "content": "question"}], + ), + response=ChatCompletion.model_validate( + { + "id": "chat", + "object": "chat.completion", + "created": 0, + "model": "policy", + "choices": [ + { + "index": index, + "finish_reason": "stop", + "message": {"role": "assistant", "content": text}, + "prompt_token_ids": [1], + "token_ids": [token_id], + "logprobs": { + "content": [ + { + "token": f"token_id:{token_id}", + "logprob": -0.1 * token_id, + "bytes": [], + "top_logprobs": [], + } + ] + }, + } + for index, (text, token_id) in enumerate( + (("left", 2), ("right", 3)) + ) + ], + } + ), + start_time=datetime(2026, 1, 1), + end_time=datetime(2026, 1, 1), + ) + group = art.TrajectoryGroup( + [art.Trajectory(exchanges=TrajectoryExchanges(chat_completions=[exchange]))] + ) + + single_exchange = exchange.model_copy( + update={ + "response": exchange.response.model_copy( + update={"choices": [exchange.response.choices[0]]} + ) + } + ) + single = art.TrajectoryGroup( + [ + art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[single_exchange]) + ) + ] + ).tokenize() + single_json = single.model_dump_json() + assert '"NaN"' in single_json + art.TokenizedTrajectoryGroup[art.TokenizedTrajectory].model_validate_json( + single_json + ) + + multi = group.tokenize(multi_history=True) + multi_json = multi.model_dump_json() + assert '"NaN"' in multi_json + art.TokenizedTrajectoryGroup[ + art.TokenizedMultiHistoryTrajectory + ].model_validate_json(multi_json) diff --git a/uv.lock b/uv.lock index 4709e1923..cf4aae62d 100644 --- a/uv.lock +++ b/uv.lock @@ -4882,7 +4882,7 @@ requires-dist = [ { name = "anthropic", specifier = ">=0.77.0" }, { name = "apex", marker = "extra == 'megatron'", git = "https://github.com/NVIDIA/apex.git?rev=25.09" }, { name = "awscli", marker = "extra == 'backend'", specifier = ">=1.38.1" }, - { name = "bitsandbytes", marker = "extra == 'backend'", specifier = ">=0.45.2" }, + { name = "bitsandbytes", marker = "extra == 'backend'", specifier = ">=0.45.2,!=0.50.0" }, { name = "causal-conv1d", marker = "python_full_version < '3.12' and platform_machine == 'x86_64' and sys_platform == 'linux' and extra == 'megatron'", specifier = "==1.6.1" }, { name = "datrie", marker = "extra == 'tinker'", specifier = ">=0.8.3" }, { name = "duckdb", marker = "extra == 'backend'", specifier = ">=1.0.0" }, diff --git a/vllm_runtime/src/art_vllm_runtime/patches.py b/vllm_runtime/src/art_vllm_runtime/patches.py index 97bc4fdb7..cef798784 100644 --- a/vllm_runtime/src/art_vllm_runtime/patches.py +++ b/vllm_runtime/src/art_vllm_runtime/patches.py @@ -110,12 +110,9 @@ def subclass_chat_completion_request() -> None: return class ChatCompletionRequest(protocol.ChatCompletionRequest): - def __init__(self, *args: object, **kwargs: object) -> None: - super().__init__(*args, **kwargs) # ty:ignore[invalid-argument-type] - self.logprobs = True - if self.top_logprobs is None: - self.top_logprobs = 0 - self.return_token_ids = True + logprobs: bool | None = True + top_logprobs: int | None = 0 + return_token_ids: bool | None = True protocol.ChatCompletionRequest = ChatCompletionRequest # ty:ignore[invalid-assignment] setattr(protocol, "_art_chat_completion_request_patched", True)