diff --git a/docs/reference/public-api.mdx b/docs/reference/public-api.mdx index 08e6825..270cb0e 100644 --- a/docs/reference/public-api.mdx +++ b/docs/reference/public-api.mdx @@ -318,7 +318,12 @@ These modules are public once their plugin package is installed: - `Proxy`, `on`, `Request`, `ClientResponse`, `Client`, `Handler`, `TunnelHandle`, `AbridgeError`, `DynamicRoutes` - `Forward`, `SessionForward` (forward to a host-side sidecar / gateway) - - `Recorder` (record tunnel request/response pairs to JSONL) + - `Recorder` (record tunnel traffic as `abridge.record.v1` JSONL) + - `CaptureLevel`, `RequestFacts`, `ResponseFacts`, `PrefixRelation`, + `PrefixTracker`, `request_facts(...)`, `response_facts(...)`, + `RECORD_SCHEMA_VERSION`, `SESSION_META_SCHEMA_VERSION` + (the record's capture ladder and derivations — see + `agentix/bridge/capture.py`) - `Sidecar`, `SidecarError`, `Command` - handler clients under `agentix.bridge.clients` (`AnthropicClient`, `OpenAIClient`, `AnthropicFromOpenAIClient`, `AnthropicToOpenAI`) diff --git a/plugins/abridge/README.md b/plugins/abridge/README.md index fe2b2d5..7c3adca 100644 --- a/plugins/abridge/README.md +++ b/plugins/abridge/README.md @@ -131,22 +131,60 @@ Claude-speaking agent in front of an OpenAI-shaped recording gateway. ### Record the tunnel traffic -`Recorder` wraps any handler client and appends one JSONL line per served -call — `{ts, path, request_id, session_id?, request, response}` — flushed -as it goes, so the file is complete up to the last call even if the host -dies mid-rollout. It exposes the wrapped client's routes and closes it on -teardown, so it drops in transparently: +`Recorder` wraps any handler client and appends one `abridge.record.v1` JSONL +line per served call, flushed as it goes, so the file is complete up to the +last call even if the host dies mid-rollout. It exposes the wrapped client's +routes and closes it on teardown, so it drops in transparently: ```python -from agentix.bridge import Proxy, Recorder +from agentix.bridge import CaptureLevel, Proxy, Recorder -proxy = Proxy(Recorder(client, "runs/rollout-42.jsonl", session_id="rollout-42")) +proxy = Proxy(Recorder( + client, "runs/rollout-42.jsonl", + session_id="rollout-42", + level=CaptureLevel.VERBATIM, # default is METADATA +)) ``` +How much lands on disk is a level, and full capture is **opt-in**: + +| level | on disk | +|---|---| +| `off` | nothing — don't install a Recorder | +| `metadata` (default) | identity, model, sampling, usage, tool **names**, per-message digests, the response block skeleton, and the prefix relation. No conversation text. | +| `verbatim` | all of the above **plus** the complete request body (`system`, `tools` with full schemas, the entire message history) and the complete structured response (`thinking` with its opaque `signature`, `text`, `tool_use`). | + +`agentix/bridge/capture.py`'s module docstring is the normative schema +document. The two fields worth knowing about up front: + +- **`prefix`** — whether this request's message list *extends* the previous + one. `"stable": false` means **the context was rewritten** (compaction, a + retry rollback, a fresh conversation on the same key), with + `divergence_index` pointing at the first message that differs. A consumer + reads that one boolean instead of parsing a harness's own compaction + boundary records. It is the Anthropic-face analogue of the token gateway's + `prefix_stable`, computed over per-message digests because this layer has + no tokenizer. Interleaved conversations (subagents, helper calls) are kept + in separate lanes keyed by the system prompt, so alternating between them + is not mistaken for a rewrite; a lane collision reports `false`, never a + false `true`. +- **`shape.content_blocks`** — one entry per block the model emitted, so + `{"type": "thinking", "chars": 0, "signature_chars": 210}` is visible even + at the metadata level. Streaming responses are recorded as the structured + `Message` (the bundled clients publish it), not as an SSE blob to re-parse. + The `request_id` in each row is the same id the transport stamps as `x-request-id` on the upstream hop (bound through a context var), so a message-level row joins a downstream token recorder's per-turn record; `session_id`, when given, tags every row with the rollout identity. +`turn_index` is monotonic per file and advances even when a row fails to +serialize, so a dropped row leaves a detectable gap; `aclose()` appends an +`abridge.session.v1` trailer, so a truncated file is distinguishable from a +closed one. + +Records never contain credentials — the tunnel carries no HTTP metadata at +all (`Request` is a path plus the decoded JSON body), so no `Authorization` +header or API key can reach one. Files are created `0600`. ## Writing your own handler @@ -293,6 +331,9 @@ More serve options: own session id, i.e. the `session_id` in its token records. The same `request_id` reaches the upstream as `x-request-id`, so message rows join the gateway's token records per call as well as per session. +* `--capture-level {off,metadata,verbatim}` (env `ABRIDGE_CAPTURE_LEVEL`, + default `metadata`) — how much of each call `--record-dir` persists. + Turning recording on never turns verbatim capture on; ask for it. `GET /_health` reports `translation_spec_sha` — one SHA-256 over the source of the Anthropic↔OpenAI transform module and both client modules @@ -315,6 +356,8 @@ agentix/bridge/ ├── proxy.py # Proxy + @on + sandbox tunnel + wire types ├── serve.py # direct mode: @on handlers as a standalone HTTP service ├── forward.py # JSON POST forwarding to a host-side service +├── capture.py # abridge.record.v1: capture levels, digests, prefix relation +├── recorder.py # Recorder: the JSONL sink for that record ├── sidecar.py # local process lifecycle + health supervision └── clients/ # bundled handler implementations ├── openai.py # OpenAIClient (openai SDK) + PLACEHOLDER_API_KEY diff --git a/plugins/abridge/ROADMAP.md b/plugins/abridge/ROADMAP.md index 9d95765..e676721 100644 --- a/plugins/abridge/ROADMAP.md +++ b/plugins/abridge/ROADMAP.md @@ -71,11 +71,13 @@ clients remain valid while sidecar gateways mature. `@on(path)` by index. Useful for offline eval reruns, RL buffer regression tests, CI-friendly assertions without burning tokens. -4. **Capture API.** Today storage is gone from the core (skipped in - the recent cleanup). Add `agentix.bridge.capture` — a small hook - that any handler can call (or a Proxy-level event subscriber) to - record full `(request, response)` pairs. Lightweight; in-memory - list with optional `JsonlSink` / `ParquetSink` overlays. +4. **Capture API.** Shipped as `agentix.bridge.capture` (the + `abridge.record.v1` shape, capture levels, and the prefix relation) + plus the `Recorder` JSONL sink. Still open: an in-memory sink for + tests, a columnar (`ParquetSink`) overlay for large corpora, and + structure recovery for an opaque SSE body relayed by a bare + `Forward` — today such a row keeps the bytes and says so rather + than reconstructing blocks the tunnel never decoded. ## Medium-term — additional bundled clients diff --git a/plugins/abridge/agentix/bridge/__init__.py b/plugins/abridge/agentix/bridge/__init__.py index 5f8ae62..cffae38 100644 --- a/plugins/abridge/agentix/bridge/__init__.py +++ b/plugins/abridge/agentix/bridge/__init__.py @@ -33,6 +33,17 @@ from __future__ import annotations +from .capture import ( + RECORD_SCHEMA_VERSION, + SESSION_META_SCHEMA_VERSION, + CaptureLevel, + PrefixRelation, + PrefixTracker, + RequestFacts, + ResponseFacts, + request_facts, + response_facts, +) from .forward import Forward, SessionForward from .proxy import ( NAMESPACE, @@ -52,21 +63,30 @@ __version__ = "0.5.0" __all__ = [ + "NAMESPACE", + "RECORD_SCHEMA_VERSION", + "SESSION_META_SCHEMA_VERSION", "AbridgeError", + "CaptureLevel", "Client", "ClientResponse", "Command", "DynamicRoutes", "Forward", "Handler", - "NAMESPACE", + "PrefixRelation", + "PrefixTracker", "Proxy", "Recorder", "Request", + "RequestFacts", + "ResponseFacts", "SessionForward", "Sidecar", "SidecarError", "TunnelHandle", "__version__", "on", + "request_facts", + "response_facts", ] diff --git a/plugins/abridge/agentix/bridge/_request_id.py b/plugins/abridge/agentix/bridge/_request_id.py index a44ee9e..a7441f1 100644 --- a/plugins/abridge/agentix/bridge/_request_id.py +++ b/plugins/abridge/agentix/bridge/_request_id.py @@ -22,12 +22,22 @@ after the lazy session create). The `Recorder` clears it before each handler call and reads it afterwards into the row's `gateway_session_id`, restoring the session-level join between caller-side rows and gateway-side records. + +`current_response_message` flows the same way and carries the STRUCTURE the +wire loses. An agent that asked for streaming gets `text/event-stream` bytes +back, so a capture layer reading only `ClientResponse` would have to re-parse +an SSE blob to find the assistant's `thinking` / `text` / `tool_use` blocks — +exactly the kind of string-scraping this capture exists to eliminate. The +Anthropic-face clients already hold the completed `Message` dict at that +point (they drain the stream and re-render it), so they publish it here and +the `Recorder` writes the object, not the blob. """ from __future__ import annotations import uuid from contextvars import ContextVar +from typing import Any current_request_id: ContextVar[str | None] = ContextVar("abridge_request_id", default=None) @@ -35,6 +45,10 @@ "abridge_upstream_session_id", default=None ) +current_response_message: ContextVar[dict[str, Any] | None] = ContextVar( + "abridge_response_message", default=None +) + def mint_request_id() -> str: return uuid.uuid4().hex @@ -46,9 +60,21 @@ def get_or_mint_request_id() -> str: return bound if bound else mint_request_id() +def publish_response_message(message: dict[str, Any]) -> None: + """Hand the completed provider-native response to an outer capture layer. + + A no-op when nothing is capturing — setting a `ContextVar` nobody reads + costs one assignment, so clients call this unconditionally rather than + branching on whether a `Recorder` happens to be installed. + """ + current_response_message.set(message) + + __all__ = [ "current_request_id", + "current_response_message", "current_upstream_session_id", "get_or_mint_request_id", "mint_request_id", + "publish_response_message", ] diff --git a/plugins/abridge/agentix/bridge/capture.py b/plugins/abridge/agentix/bridge/capture.py new file mode 100644 index 0000000..a27c78c --- /dev/null +++ b/plugins/abridge/agentix/bridge/capture.py @@ -0,0 +1,537 @@ +"""`abridge.record.v1` — the normative capture shape for the Anthropic face. + +The OTel GenAI span the bundled clients stamp (`_genai_span.py`) records a +call as `(role, content-flattened-to-text)` pairs plus token counts. That is +enough to *watch* a rollout and not nearly enough to *reconstruct* one: the +tool schemas are gone, `thinking` signatures are gone, a `tool_use` block +becomes a printf'd string, and nothing says whether the context was rewritten +between two calls. This module defines the record that closes that gap, and +`recorder.py` writes it. + +**This module docstring is the normative schema document** — downstream +consumers adapt to the shape defined here (the same contract stance the token +gateway's `tito.record.v1` takes). + +## Capture levels + +Verbatim bodies are large and contain whatever the user typed, so full +capture is **opt-in** and the ladder is strictly additive: + +`off` + Nothing — no recorder is installed. +`metadata` + Identity, model, sampling, usage, tool NAMES, per-message digests, the + response block skeleton, and the prefix relation. No conversation text. +`verbatim` + Everything above **plus** the complete request body (`system`, `tools` + with full schemas, the entire message history) and the complete + structured response (`thinking` + its opaque `signature`, `text`, + `tool_use`). + +`metadata` is the default: an operator who turns recording on gets the join +keys and the rewrite signal without their users' prompts landing on disk. +Nothing silently promotes a configured level. + +`shape`, `sampling`, `conversation_key`, and `prefix` are derived from the +Anthropic Messages body, so they appear only on `MESSAGES_PATH` rows. A +`Recorder` wrapping some other face still captures that face's bodies at the +verbatim level; it just reports no Anthropic-face derivations for them. + +## The record + +One JSON line per served call:: + + { + "schema_version": "abridge.record.v1", + "session_id": "...", // OPTIONAL: only when the Recorder has one + "gateway_session_id": "...", // OPTIONAL: only when the transport published one + "request_id": "<32 hex>", // == the x-request-id stamped upstream + "turn_index": 0, // monotonic per file; gaps mark dropped rows + "ts": 1750000000.0, + "path": "/v1/messages", + "capture_level": "verbatim", // the level ACTUALLY applied to this row + "model": "claude-sonnet-4-5", // the model field the agent sent + "stream": true, // what the agent asked for + "sampling": {"max_tokens": 8192, "thinking": {...}, ...}, + "conversation_key": "<16 hex>", // lane id — see "Prefix relation" + "shape": { + "system_digest": "<16 hex>", // null when the request carried no system + "tools_digest": "<16 hex>", // over the FULL schemas, not the names + "tool_names": ["Bash", "Agent", ...], + "messages": 42, + "message_digests": ["<16 hex>", ...], // one per message, in order + "content_blocks": [ // what the model emitted + {"type": "thinking", "chars": 0, "signature_chars": 210}, + {"type": "text", "chars": 412}, + {"type": "tool_use", "name": "Bash", "input_chars": 88} + ], + "response_digest": "<16 hex>" + }, + "prefix": { + "stable": false, // false ⇒ THE CONTEXT WAS REWRITTEN + "divergence_index": 3, // first message index that differs + "common_prefix": 3, + "previous_request_id": "...", + "previous_messages": 41, + "system_changed": false, + "tools_changed": false, + "assistant_echo": "modified" // "verbatim" | "modified" | "absent" | null + }, + "status_code": 200, + "media_type": "text/event-stream", + "stop_reason": "tool_use", + "usage": {"input_tokens": 30112, "output_tokens": 214, ...}, + + // verbatim level only (and never on `VERBATIM_EXCLUDED_PATHS`): + "request": { ...the complete request body, exactly as the agent sent it... }, + "response": { ...the complete Anthropic Message, incl. thinking signatures... }, + "response_body": "event: ...", // only when no structured message exists + + // instead of stop_reason/usage/response when the handler raised: + "error": "AbridgeError: upstream exploded" + } + +Closing the recorder appends one final line:: + + {"schema_version": "abridge.session.v1", "session_id": ..., "turns": N, + "capture_level": "verbatim", "ts": ...} + +A file without that line was truncated (the process died mid-rollout); a +`turn_index` gap marks a row that failed to serialize or write. + +## Prefix relation — the compaction signal + +The load-bearing field. **When request N+1's message list is not an extension +of request N's, the context was rewritten**, and `prefix.stable` is `false`. +That single fact subsumes compaction detection: a consumer never has to parse +a harness's own boundary records, follow a `parentUuid: null`, or scan a +summary string. It is the Anthropic-face analogue of the token gateway's +`prefix_stable` (which asks the same question of token ids), computed here +over per-message digests because this layer has no tokenizer and must not +invent one. + +One API key multiplexes several logical conversations (helper calls, +subagents, reruns), so a single global "previous request" would report a +rewrite on every alternation and the signal would be worthless. Requests are +therefore attributed to a **lane** keyed by the canonicalized system prompt +(`conversation_key`) — the same demux idea `serve.py` already uses for the +token gateway, minus the first user message, because compaction *replaces* +the first user message and keying on it would hide exactly the event we are +trying to surface. Two unrelated conversations that share a system prompt +collide into one lane and read as `stable: false`; a lane evicted past +`max_lanes` reads as a fresh conversation (`previous_request_id: null`). +Both are the safe direction — the computation is deliberately biased so that +a spurious "rewritten" costs a consumer one split, where a spurious "stable" +would splice two unrelated contexts into one training sample. + +`assistant_echo` answers the follow-up question — whether the assistant turn +the model produced came back verbatim in the next request's history +(`"verbatim"`), came back altered (`"modified"`, e.g. the harness dropped the +`thinking` blocks), or never came back at all (`"absent"`). `null` when there +is no previous turn to compare against or the prefix already diverged. + +## Relationship to the token gateway's `tito.record.v1` + +The two records deliberately share an **envelope** and diverge on the +**payload**. Shared field names and semantics: `schema_version`, `session_id`, +`request_id`, `turn_index` (monotonic per file, advancing through failures so +a dropped row leaves a detectable gap), `ts`, `model`, `sampling` (a +whitelist, not a passthrough), a closing session line, strict JSON +(`allow_nan=False`), log-and-serve on write failure, and the meaning of the +prefix fact. One parser reads both streams and one join key (`request_id`) +lines up a row here with a turn there. + +They do not share the payload, and forcing them to would be a lie. The token +gateway owns a tokenizer, so it records `prompt_token_ids`, +`completion_logprobs`, and the segment tiling a trainer needs, and asks the +prefix question of token ids. This layer sits in front of a provider it does +not tokenize for, so it records provider-native JSON and asks the prefix +question of message digests. A record here plus the gateway's record for the +same `request_id` is the complete picture; neither one can be rewritten into +the other. + +## What this layer cannot see + +The tunnel carries no HTTP metadata at all (`proxy.Request` is `path` + +decoded JSON body; `ClientResponse` is bytes + media type + status), so no +`Authorization` header, no `x-api-key`, and no caller credential can reach a +record — not by policy but by construction. Record files are created `0600`. + +What is NOT filtered: a secret the user typed into a prompt is conversation +content, and at the verbatim level conversation content is exactly what gets +written. `verbatim` means verbatim; a redaction pass over message bodies +would quietly break the one guarantee this record exists to provide. The +level gate, the file mode, and the operator's choice of record directory are +the controls. +""" + +from __future__ import annotations + +import hashlib +import json +from collections import OrderedDict +from collections.abc import Mapping, Sequence +from dataclasses import dataclass +from enum import StrEnum +from typing import Any + +RECORD_SCHEMA_VERSION = "abridge.record.v1" +SESSION_META_SCHEMA_VERSION = "abridge.session.v1" + +# The Anthropic Messages face. `request_facts` / `response_facts` understand +# this body shape and no other, so only rows on this path carry `shape`, +# `sampling`, `conversation_key`, and `prefix`. A `Recorder` wrapping an +# OpenAI-face client still records that client's bodies verbatim; it just +# reports no Anthropic-face derivations for them. +MESSAGES_PATH = "/v1/messages" + +# Excluded from verbatim body capture even at the verbatim level: a +# token-count call re-sends a duplicate of the adjacent history and produces +# no completion, so its body roughly doubles the file and reconstructs +# nothing. Those rows say `"capture_level": "metadata"` rather than looking +# like verbatim rows that lost their body. +VERBATIM_EXCLUDED_PATHS = ("/v1/messages/count_tokens",) + +# Request-control parameters lifted verbatim into `sampling`. Whitelist, not +# passthrough: the body also carries `system` / `messages` / `tools`, which are +# conversation content and belong behind the verbatim level. +SAMPLING_KEYS = ( + "max_tokens", + "temperature", + "top_p", + "top_k", + "stop_sequences", + "thinking", + "tool_choice", + "service_tier", +) + +_DIGEST_CHARS = 16 + + +class CaptureLevel(StrEnum): + """How much of a call lands on disk. Strictly additive, opt-in upward.""" + + OFF = "off" + METADATA = "metadata" + VERBATIM = "verbatim" + + +def canonical_text(value: Any) -> str: + """Flatten Anthropic content (str or block list) to identity text. + + Text blocks contribute their text — `cache_control` and other decorations + that do not change what the model reads are ignored — and every other + block its sorted-JSON form. + """ + if value is None: + return "" + if isinstance(value, str): + return value + if isinstance(value, list): + parts: list[str] = [] + for block in value: + if isinstance(block, str): + parts.append(block) + elif isinstance(block, dict) and block.get("type") == "text": + parts.append(str(block.get("text", ""))) + else: + parts.append(json.dumps(block, sort_keys=True, ensure_ascii=False, default=repr)) + return "\n".join(parts) + return json.dumps(value, sort_keys=True, ensure_ascii=False, default=repr) + + +def digest(value: Any) -> str: + """Stable short digest of a JSON-able value (sorted keys, no whitespace). + + Truncated to 64 bits: it only ever compares values inside one session's + record stream, where a collision would have to hit two messages at the + same index of one conversation. + """ + blob = json.dumps(value, sort_keys=True, ensure_ascii=False, separators=(",", ":"), default=repr) + return hashlib.sha256(blob.encode()).hexdigest()[:_DIGEST_CHARS] + + +@dataclass(frozen=True, slots=True) +class RequestFacts: + """Everything the `metadata` level knows about one Messages request.""" + + model: str | None + stream: bool + sampling: dict[str, Any] + conversation_key: str + system_digest: str | None + tools_digest: str | None + tool_names: tuple[str, ...] + message_digests: tuple[str, ...] + + def to_dict(self) -> dict[str, Any]: + return { + "system_digest": self.system_digest, + "tools_digest": self.tools_digest, + "tool_names": list(self.tool_names), + "messages": len(self.message_digests), + "message_digests": list(self.message_digests), + } + + +def request_facts(body: Mapping[str, Any]) -> RequestFacts: + """Derive the metadata-level view of an Anthropic Messages request body. + + `tools_digest` covers the full schemas, so a consumer can tell "the tool + set changed" from "the same tools were re-sent" without the schemas + themselves; `tool_names` is the list the model actually saw on the wire — + which is not necessarily the list a harness reports about itself. + """ + system = body.get("system") + tools = body.get("tools") + messages = body.get("messages") + names: list[str] = [] + if isinstance(tools, list): + for tool in tools: + if isinstance(tool, dict): + name = tool.get("name") + if isinstance(name, str): + names.append(name) + return RequestFacts( + model=body.get("model") if isinstance(body.get("model"), str) else None, + stream=bool(body.get("stream", False)), + sampling={key: body[key] for key in SAMPLING_KEYS if body.get(key) is not None}, + conversation_key=digest(canonical_text(system)), + system_digest=None if system is None else digest(system), + tools_digest=None if tools is None else digest(tools), + tool_names=tuple(names), + message_digests=tuple(digest(m) for m in messages) if isinstance(messages, list) else (), + ) + + +@dataclass(frozen=True, slots=True) +class ResponseFacts: + """The metadata-level view of one Anthropic Message the model produced.""" + + stop_reason: str | None + usage: dict[str, Any] | None + content_blocks: tuple[dict[str, Any], ...] | None + response_digest: str | None + assistant_digest: str | None + + def to_dict(self) -> dict[str, Any]: + return { + "content_blocks": None if self.content_blocks is None else [dict(b) for b in self.content_blocks], + "response_digest": self.response_digest, + } + + +EMPTY_RESPONSE_FACTS = ResponseFacts( + stop_reason=None, usage=None, content_blocks=None, response_digest=None, assistant_digest=None +) + + +def response_facts(message: Mapping[str, Any] | None) -> ResponseFacts: + """Derive the metadata-level view of a structured Anthropic Message. + + `content_blocks` records one entry per emitted block. A `thinking` block + on providers that return redacted reasoning arrives with empty text and a + populated opaque `signature`, so both lengths are recorded separately — + `chars: 0, signature_chars: 210` is a meaningful, checkable shape, and the + signature itself is preserved verbatim one level up. + """ + if message is None: + return EMPTY_RESPONSE_FACTS + content = message.get("content") + blocks: list[dict[str, Any]] | None = None + if isinstance(content, list): + blocks = [_block_shape(block) for block in content] + usage = message.get("usage") + stop_reason = message.get("stop_reason") + return ResponseFacts( + stop_reason=stop_reason if isinstance(stop_reason, str) else None, + usage=dict(usage) if isinstance(usage, Mapping) else None, + content_blocks=None if blocks is None else tuple(blocks), + response_digest=digest(message), + # Keyed exactly like a history message so it can be compared against + # the assistant turn the agent echoes back on the next request. + assistant_digest=digest({"role": "assistant", "content": content}), + ) + + +def _block_shape(block: Any) -> dict[str, Any]: + if not isinstance(block, dict): + return {"type": None, "chars": len(str(block))} + block_type = block.get("type") + if block_type == "text": + return {"type": "text", "chars": len(str(block.get("text", "")))} + if block_type == "thinking": + return { + "type": "thinking", + "chars": len(str(block.get("thinking", ""))), + "signature_chars": len(str(block.get("signature", ""))), + } + if block_type == "redacted_thinking": + return {"type": "redacted_thinking", "data_chars": len(str(block.get("data", "")))} + if block_type == "tool_use": + name = block.get("name") + return { + "type": "tool_use", + "name": name if isinstance(name, str) else None, + "input_chars": len(json.dumps(block.get("input") or {}, ensure_ascii=False, default=repr)), + } + return { + "type": block_type if isinstance(block_type, str) else None, + "chars": len(json.dumps(block, ensure_ascii=False, default=repr)), + } + + +@dataclass(frozen=True, slots=True) +class PrefixRelation: + """How this request's message list relates to the lane's previous one.""" + + stable: bool + divergence_index: int | None + common_prefix: int + previous_request_id: str | None + previous_messages: int | None + system_changed: bool + tools_changed: bool + assistant_echo: str | None + + def to_dict(self) -> dict[str, Any]: + return { + "stable": self.stable, + "divergence_index": self.divergence_index, + "common_prefix": self.common_prefix, + "previous_request_id": self.previous_request_id, + "previous_messages": self.previous_messages, + "system_changed": self.system_changed, + "tools_changed": self.tools_changed, + "assistant_echo": self.assistant_echo, + } + + +# The first request seen on a lane: nothing to extend, so nothing is broken. +# Mirrors the token gateway's `prefix_stable = not last or ...`. +FIRST_TURN = PrefixRelation( + stable=True, + divergence_index=None, + common_prefix=0, + previous_request_id=None, + previous_messages=None, + system_changed=False, + tools_changed=False, + assistant_echo=None, +) + + +def common_prefix_length(previous: Sequence[str], current: Sequence[str]) -> int: + """How many leading entries the two digest sequences share.""" + count = 0 + for before, after in zip(previous, current): + if before != after: + break + count += 1 + return count + + +@dataclass(slots=True) +class _Lane: + request_id: str + message_digests: tuple[str, ...] + system_digest: str | None + tools_digest: str | None + assistant_digest: str | None = None + + +class PrefixTracker: + """Per-conversation "is this an extension of the last request" bookkeeping. + + `begin` returns the relation for a request and installs it as its lane's + new baseline; `finish` attaches the assistant turn that request produced, + so the next `begin` can also report whether that turn was echoed back + faithfully. Lanes are keyed by `RequestFacts.conversation_key` and + LRU-bounded at `max_lanes` — an evicted lane's next request reads as a + fresh conversation, which is why the bound should exceed the number of + conversations one agent multiplexes (main loop + concurrent subagents). + + Not concurrency-safe in the strict sense: two calls racing on ONE lane + interleave their baselines. The tracker is used from the recorder's own + handler coroutine, and the failure direction is `stable: false`, so a race + costs a spurious split rather than a bad splice. + """ + + def __init__(self, *, max_lanes: int = 16) -> None: + if max_lanes < 1: + raise ValueError(f"max_lanes must be >= 1, got {max_lanes!r}") + self._max_lanes = max_lanes + self._lanes: OrderedDict[str, _Lane] = OrderedDict() + + def begin(self, *, request_id: str, facts: RequestFacts) -> PrefixRelation: + key = facts.conversation_key + lane = self._lanes.get(key) + relation = FIRST_TURN if lane is None else self._relate(lane, facts) + self._lanes[key] = _Lane( + request_id=request_id, + message_digests=facts.message_digests, + system_digest=facts.system_digest, + tools_digest=facts.tools_digest, + ) + self._lanes.move_to_end(key) + while len(self._lanes) > self._max_lanes: + self._lanes.popitem(last=False) + return relation + + @staticmethod + def _relate(lane: _Lane, facts: RequestFacts) -> PrefixRelation: + previous = lane.message_digests + current = facts.message_digests + shared = common_prefix_length(previous, current) + # An extension keeps every previous message, in order, and adds to it. + # A shorter list that agrees as far as it goes is a truncation, i.e. + # still a rewrite — `len(current) >= len(previous)` catches that. + stable = shared == len(previous) and len(current) >= len(previous) + echo: str | None = None + if stable and lane.assistant_digest is not None: + if len(current) > len(previous): + echo = "verbatim" if current[len(previous)] == lane.assistant_digest else "modified" + else: + echo = "absent" + return PrefixRelation( + stable=stable, + divergence_index=None if stable else shared, + common_prefix=shared, + previous_request_id=lane.request_id, + previous_messages=len(previous), + system_changed=lane.system_digest != facts.system_digest, + tools_changed=lane.tools_digest != facts.tools_digest, + assistant_echo=echo, + ) + + def finish(self, *, request_id: str, conversation_key: str, assistant_digest: str | None) -> None: + """Attach the assistant turn `request_id` produced to its lane. + + Ignored when another call has already taken the lane over (an + interleaved `begin`), so a race never mislabels somebody else's turn. + """ + lane = self._lanes.get(conversation_key) + if lane is not None and lane.request_id == request_id: + lane.assistant_digest = assistant_digest + + +__all__ = [ + "EMPTY_RESPONSE_FACTS", + "FIRST_TURN", + "MESSAGES_PATH", + "RECORD_SCHEMA_VERSION", + "SAMPLING_KEYS", + "SESSION_META_SCHEMA_VERSION", + "VERBATIM_EXCLUDED_PATHS", + "CaptureLevel", + "PrefixRelation", + "PrefixTracker", + "RequestFacts", + "ResponseFacts", + "canonical_text", + "common_prefix_length", + "digest", + "request_facts", + "response_facts", +] diff --git a/plugins/abridge/agentix/bridge/clients/anthropic.py b/plugins/abridge/agentix/bridge/clients/anthropic.py index 8140b56..6361e14 100644 --- a/plugins/abridge/agentix/bridge/clients/anthropic.py +++ b/plugins/abridge/agentix/bridge/clients/anthropic.py @@ -27,6 +27,7 @@ from agentix.utils import trace +from .._request_id import get_or_mint_request_id, publish_response_message from ..proxy import ( AbridgeError, ClientResponse, @@ -103,7 +104,9 @@ async def messages(self, request: Request) -> ClientResponse: body = dict(request.body) if self._model: body["model"] = self._model - record_id = uuid.uuid4().hex + # Reuses the id a wrapping capture layer (Recorder) bound for this + # call, so its JSONL row and the upstream header share one id. + record_id = get_or_mint_request_id() extra_headers = { "x-session-id": self.session_id, "x-request-id": record_id, @@ -127,6 +130,10 @@ async def messages(self, request: Request) -> ClientResponse: message = await stream_handle.get_final_message() response_dict = message.model_dump(exclude_none=False) populate_anthropic_span(request=request.body, response=response_dict) + # The wire the agent gets is an SSE blob; hand the capture + # layer the completed Message so `thinking` signatures and + # `tool_use` inputs are recorded as structure, not text. + publish_response_message(response_dict) return ClientResponse.sse(_anthropic_sse_from_message(response_dict)) message = await self._client.messages.create( @@ -137,6 +144,7 @@ async def messages(self, request: Request) -> ClientResponse: raise AbridgeError(f"anthropic: {exc}", status_code=status) from exc response_dict = message.model_dump(exclude_none=False) populate_anthropic_span(request=request.body, response=response_dict) + publish_response_message(response_dict) return ClientResponse.json(response_dict) @on("/v1/messages/count_tokens") diff --git a/plugins/abridge/agentix/bridge/clients/anthropic_from_openai.py b/plugins/abridge/agentix/bridge/clients/anthropic_from_openai.py index 3b67508..d336f7e 100644 --- a/plugins/abridge/agentix/bridge/clients/anthropic_from_openai.py +++ b/plugins/abridge/agentix/bridge/clients/anthropic_from_openai.py @@ -19,7 +19,7 @@ from agentix.utils import trace -from .._request_id import get_or_mint_request_id +from .._request_id import get_or_mint_request_id, publish_response_message from ..proxy import ( AbridgeError, ClientResponse, @@ -127,6 +127,9 @@ async def messages(self, request: Request) -> ClientResponse: openai_resp, response_model=str(request.body.get("model") or "") ) populate_anthropic_span(request=request.body, response=anthropic_resp) + # Agent-facing shape, so the capture layer records the Anthropic + # Message the agent actually saw — structure, not an SSE blob. + publish_response_message(anthropic_resp) if request.body.get("stream"): return ClientResponse.sse(anthropic_sse(anthropic_resp)) return ClientResponse.json(anthropic_resp) diff --git a/plugins/abridge/agentix/bridge/clients/anthropic_to_openai.py b/plugins/abridge/agentix/bridge/clients/anthropic_to_openai.py index 5d5a42f..fb21708 100644 --- a/plugins/abridge/agentix/bridge/clients/anthropic_to_openai.py +++ b/plugins/abridge/agentix/bridge/clients/anthropic_to_openai.py @@ -32,6 +32,7 @@ from agentix.utils import trace +from .._request_id import publish_response_message from ..proxy import ClientResponse, Handler, Request, TunnelHandle, _AsyncCloseable, on from ._anthropic_transforms import ( anthropic_messages_to_openai, @@ -98,6 +99,9 @@ async def messages(self, request: Request) -> ClientResponse: openai_resp, response_model=str(request.body.get("model") or "") ) populate_anthropic_span(request=request.body, response=anthropic_resp) + # Agent-facing shape, so the capture layer records the Anthropic + # Message the agent actually saw — structure, not an SSE blob. + publish_response_message(anthropic_resp) if request.body.get("stream"): return ClientResponse.sse(anthropic_sse(anthropic_resp)) return ClientResponse.json(anthropic_resp) diff --git a/plugins/abridge/agentix/bridge/recorder.py b/plugins/abridge/agentix/bridge/recorder.py index b34e054..0b6a2af 100644 --- a/plugins/abridge/agentix/bridge/recorder.py +++ b/plugins/abridge/agentix/bridge/recorder.py @@ -1,57 +1,87 @@ """`Recorder` — capture rollout traffic at the tunnel, one JSONL line per call. -The tunnel is the one place every LLM call an agent makes passes through, -so it is the natural recording point for rollout data collection: wrap any -handler client in `Recorder(client, path)` and hand the wrapper to +The tunnel is the one place every LLM call an agent makes passes through, so +it is the natural recording point for rollout data collection: wrap any +handler client in `Recorder(client, path, level=...)` and hand the wrapper to `Proxy(...)` — neither the agent nor the upstream can tell the difference. -Each served request appends one line:: - - {"ts": ..., "path": "/v1/messages", "request_id": "<32 hex>", - "session_id": ..., # only when the Recorder has one - "gateway_session_id": ..., # only when the transport published one - "request": {...}, - "response": {"status_code": 200, "media_type": "...", "body": ...}} - -`request_id` is minted per call and bound on the `current_request_id` -context var for the duration of the handler, so the transport layer -(`Forward` / the SDK clients) stamps the SAME id as `x-request-id` on the -upstream hop — a downstream token recorder's per-turn record and this row -join on it. `session_id`, when given, identifies the rollout the wrapped -client serves (pass the same value as the client's session identity). -`gateway_session_id` is read back from the transport after the call (via -`current_upstream_session_id`): when the downstream is a session-scoped -gateway (`SessionForward`), it is the gateway's OWN session id — i.e. the -`session_id` in the gateway's token records — restoring the session-level -join that the caller-side hash alone cannot provide. Without these keys, -rows from a retried call (e.g. an agent retry after a tunnel 504 produced an -orphan success row) are only deduplicable by request-body equality. - -A handler that raises records `{"error": ...}` instead of `"response"` and -re-raises — a failed call is signal, not something to lose. JSON bodies are -recorded as objects; anything else (e.g. a pre-rendered SSE blob) as text. +`capture.py` defines the row shape (`abridge.record.v1`) and the capture +ladder; this module is the sink that assembles, orders, and persists it. +Read that module's docstring for the schema. The short version: + + * `CaptureLevel.METADATA` (the default) writes identity, model, sampling, + usage, tool names, per-message digests, the response block skeleton, and + the prefix relation — no conversation text. + * `CaptureLevel.VERBATIM` adds the complete request body (system, full tool + schemas, the entire message history) and the complete structured response + (`thinking` with its opaque `signature`, `text`, `tool_use`). + +Full capture is opt-in and a configured level is never promoted. + +Identity and joins. `request_id` is minted per call and bound on the +`current_request_id` context var for the duration of the handler, so the +transport layer (`Forward` / the SDK clients) stamps the SAME id as +`x-request-id` on the upstream hop — a downstream token recorder's per-turn +record and this row join on it. `session_id`, when given, identifies the +rollout the wrapped client serves. `gateway_session_id` is read back from the +transport after the call (via `current_upstream_session_id`): when the +downstream is a session-scoped gateway (`SessionForward`), it is the +gateway's OWN session id — i.e. the `session_id` in the gateway's token +records — restoring the session-level join that the caller-side hash alone +cannot provide. + +Structure over blobs. An agent that asked for streaming gets `text/event- +stream` bytes back, so the clients publish the completed `Message` dict on +`current_response_message` and the row carries the object. When nothing +published one (a raw pass-through `Forward` relaying somebody else's SSE), +`response` is absent, `response_body` holds the bytes as text, and the +skeleton fields are `null` — the boundary is recorded, never guessed at. + +A handler that raises records `{"error": ...}` instead of the response fields +and re-raises — a failed call is signal, not something to lose. Handlers run on the event loop, so appends never interleave; each line is -flushed as it is written so the file is complete up to the last call even -if the process dies mid-rollout. The file opens lazily on the first record, -so a Recorder that never serves (e.g. a route-enumeration probe) leaves no -empty file behind. Capture is log-and-serve: a failed row write (disk full, -unencodable text) is logged and the agent's call still succeeds — matching -the token-recording gateway's policy, so the two capture layers never -disagree about whether a turn happened. After `aclose()` a straggler -in-flight call's row is dropped (logged), never written to a resurrected -file handle. +flushed as it is written so the file is complete up to the last call even if +the process dies mid-rollout. The file opens lazily (mode `0600`) on the +first record, so a Recorder that never serves (e.g. a route-enumeration +probe) leaves no empty file behind, and `aclose()` appends one +`abridge.session.v1` line so a truncated file is distinguishable from a +closed one. `turn_index` advances even when a line fails, so any dropped row +leaves a detectable gap. Capture is log-and-serve: a failed row write (disk +full, unencodable text) is logged and the agent's call still succeeds — +matching the token-recording gateway's policy, so the two capture layers +never disagree about whether a turn happened. After `aclose()` a straggler +in-flight call's row is dropped (logged), never written to a resurrected file +handle. """ from __future__ import annotations +import copy import json import logging +import os import time from pathlib import Path from typing import IO, Any -from ._request_id import current_request_id, current_upstream_session_id, mint_request_id +from ._request_id import ( + current_request_id, + current_response_message, + current_upstream_session_id, + mint_request_id, +) +from .capture import ( + MESSAGES_PATH, + RECORD_SCHEMA_VERSION, + SESSION_META_SCHEMA_VERSION, + VERBATIM_EXCLUDED_PATHS, + CaptureLevel, + PrefixTracker, + RequestFacts, + request_facts, + response_facts, +) from .proxy import ClientResponse, Handler, Request, _collect_handlers logger = logging.getLogger(__name__) @@ -64,33 +94,72 @@ class Recorder: dynamic-route seam), delegates `environ(...)`, and closes both the inner client and the record file on `aclose()` — so `Proxy.stop()` tears the whole stack down once, as usual. + + `level` picks how much lands on disk (see `capture.CaptureLevel`); it + defaults to `METADATA` because verbatim bodies carry whatever the user + typed. `max_lanes` bounds the prefix tracker's per-conversation state — + raise it above the number of conversations the agent multiplexes at once + (its main loop plus concurrent subagents). """ - def __init__(self, client: Any, path: str | Path, *, session_id: str | None = None) -> None: + def __init__( + self, + client: Any, + path: str | Path, + *, + session_id: str | None = None, + level: CaptureLevel = CaptureLevel.METADATA, + max_lanes: int = 16, + ) -> None: + if CaptureLevel(level) is CaptureLevel.OFF: + raise ValueError( + "Recorder cannot be built at CaptureLevel.OFF — 'off' means no recorder at all; " + "skip the wrapper instead of installing one that writes nothing" + ) self._client = client self._path = Path(path) self._path.parent.mkdir(parents=True, exist_ok=True) self._session_id = session_id + self._level = CaptureLevel(level) + self._prefix = PrefixTracker(max_lanes=max_lanes) self._file: IO[str] | None = None self._closed = False + self._turns = 0 + + @property + def level(self) -> CaptureLevel: + return self._level def abridge_routes(self) -> dict[str, Handler]: return {path: self._recording(path, handler) for path, handler in _collect_handlers(self._client).items()} def _recording(self, path: str, handler: Handler) -> Handler: + anthropic_face = path == MESSAGES_PATH + verbatim = self._level is CaptureLevel.VERBATIM and path not in VERBATIM_EXCLUDED_PATHS + async def record(request: Request) -> ClientResponse: # Reuse an id bound by an even-outer layer; otherwise mint here. # Binding it makes the transport's upstream `x-request-id` equal # this row's `request_id`. request_id = current_request_id.get() or mint_request_id() - line: dict[str, Any] = {"ts": time.time(), "path": path, "request_id": request_id} - if self._session_id is not None: - line["session_id"] = self._session_id - line["request"] = request.body + # Derived BEFORE the handler runs: these digests are the ground + # truth for what the agent sent, and the row is only serialized + # afterwards. Same reason the verbatim body is deep-copied below. + facts = request_facts(request.body) if anthropic_face else None + line = self._open_line(path=path, request_id=request_id, facts=facts, verbatim=verbatim) + if facts is not None: + line["prefix"] = self._prefix.begin(request_id=request_id, facts=facts).to_dict() + if verbatim: + # Deep-copied because the row is serialized after the handler + # ran: a handler that rewrote the body in place would + # otherwise be recorded as what the agent sent. + line["request"] = copy.deepcopy(request.body) + rid_token = current_request_id.set(request_id) # Cleared per call so a value published by a PREVIOUS call on # this task never leaks into an unrelated row. upstream_token = current_upstream_session_id.set(None) + message_token = current_response_message.set(None) try: response = await handler(request) except BaseException as exc: @@ -102,18 +171,87 @@ async def record(request: Request) -> ClientResponse: current_request_id.reset(rid_token) gateway_session_id = current_upstream_session_id.get() current_upstream_session_id.reset(upstream_token) + message = current_response_message.get() + current_response_message.reset(message_token) if gateway_session_id is not None: line["gateway_session_id"] = gateway_session_id - line["response"] = { - "status_code": response.status_code, - "media_type": response.media_type, - "body": _decode_body(response), - } + self._stamp_response( + line, response, message, facts=facts, request_id=request_id, verbatim=verbatim + ) self._write(line) return response return record + def _open_line( + self, *, path: str, request_id: str, facts: RequestFacts | None, verbatim: bool + ) -> dict[str, Any]: + applied = CaptureLevel.VERBATIM if verbatim else CaptureLevel.METADATA + line: dict[str, Any] = {"schema_version": RECORD_SCHEMA_VERSION} + if self._session_id is not None: + line["session_id"] = self._session_id + line.update( + { + "request_id": request_id, + # Reserved here to fix the key's position; the value is + # assigned at write time so two concurrent calls can't be + # handed the same index. + "turn_index": -1, + "ts": time.time(), + "path": path, + "capture_level": applied.value, + } + ) + if facts is not None: + line.update( + { + "model": facts.model, + "stream": facts.stream, + "sampling": dict(facts.sampling), + "conversation_key": facts.conversation_key, + "shape": facts.to_dict(), + } + ) + return line + + def _stamp_response( + self, + line: dict[str, Any], + response: ClientResponse, + message: dict[str, Any] | None, + *, + facts: RequestFacts | None, + request_id: str, + verbatim: bool, + ) -> None: + line["status_code"] = response.status_code + line["media_type"] = response.media_type + if message is None and response.media_type == "application/json": + # No client published structure, but a JSON body on this face IS + # the Message — decode it rather than treat it as an opaque blob. + message = _decode_json_object(response) + if facts is not None: + derived = response_facts(message) + self._prefix.finish( + request_id=request_id, + conversation_key=facts.conversation_key, + assistant_digest=derived.assistant_digest, + ) + shape = line.get("shape") + if isinstance(shape, dict): + shape.update(derived.to_dict()) + line["stop_reason"] = derived.stop_reason + line["usage"] = derived.usage + if not verbatim: + return + if message is not None: + line["response"] = message + else: + # An opaque body (somebody else's SSE relayed by a bare Forward). + # Recorded as text so nothing is lost, and flagged by the absence + # of `response` so a consumer never mistakes it for structure. + line["response_body"] = response.body.decode("utf-8", "replace") + @staticmethod def _stamp_gateway_session(line: dict[str, Any]) -> None: gateway_session_id = current_upstream_session_id.get() @@ -123,20 +261,35 @@ def _stamp_gateway_session(line: dict[str, Any]) -> None: def _write(self, line: dict[str, Any]) -> None: # Log-and-serve, mirroring the token-recording gateway's policy: the # upstream call already succeeded (or its error is being re-raised), - # so a capture failure must not turn it into a wire error. + # so a capture failure must not turn it into a wire error. The + # turn_index still advances, leaving a detectable gap. + line["turn_index"] = self._turns try: - if self._closed: - # A straggler dispatch outlived aclose(): the file is closed - # for good — dropping the row (loudly) beats resurrecting a - # file handle nobody will ever close. - logger.warning("abridge recorder: dropping row for %s — recorder is closed", self._path) - return - if self._file is None or self._file.closed: - self._file = self._path.open("a", encoding="utf-8") - self._file.write(json.dumps(line, ensure_ascii=False, default=repr) + "\n") - self._file.flush() + self._append(line) except Exception: # noqa: BLE001 - capture must never fail the served call logger.exception("abridge recorder: failed to append to %s — row NOT persisted", self._path) + finally: + self._turns += 1 + + def _append(self, line: dict[str, Any]) -> None: + if self._closed: + # A straggler dispatch outlived aclose(): the file is closed for + # good — dropping the row (loudly) beats resurrecting a file + # handle nobody will ever close. + logger.warning("abridge recorder: dropping row for %s — recorder is closed", self._path) + return + handle = self._file + if handle is None or handle.closed: + # 0600 at creation: a verbatim row holds whatever the user typed, + # so the file must not be readable by other accounts on the box. + fd = os.open(self._path, os.O_WRONLY | os.O_CREAT | os.O_APPEND, 0o600) + handle = os.fdopen(fd, "a", encoding="utf-8") + self._file = handle + # Strict JSON: no NaN/Infinity literals, matching the token gateway's + # record stream, so one parser reads both. `default=repr` covers an + # exotic value a provider slipped into an otherwise JSON body. + handle.write(json.dumps(line, ensure_ascii=False, allow_nan=False, default=repr) + "\n") + handle.flush() def environ(self, handle: Any) -> dict[str, str]: return self._client.environ(handle) @@ -147,18 +300,36 @@ async def aclose(self) -> None: if aclose is not None: await aclose() finally: - self._closed = True - if self._file is not None: - self._file.close() + self._finalize() + def _finalize(self) -> None: + """Append the session-metadata line and close the file. -def _decode_body(response: ClientResponse) -> Any: - if response.media_type == "application/json": - try: - return json.loads(response.body) - except (ValueError, UnicodeDecodeError): - pass - return response.body.decode("utf-8", "replace") + Skipped entirely for a recorder that never wrote a row, so a + route-enumeration probe still leaves no file behind. + """ + if self._closed: + return + if self._turns: + meta: dict[str, Any] = {"schema_version": SESSION_META_SCHEMA_VERSION} + if self._session_id is not None: + meta["session_id"] = self._session_id + meta.update({"turns": self._turns, "capture_level": self._level.value, "ts": time.time()}) + try: + self._append(meta) + except Exception: # noqa: BLE001 - a missing trailer is not worth failing teardown + logger.exception("abridge recorder: failed to finalize %s", self._path) + self._closed = True + if self._file is not None: + self._file.close() + + +def _decode_json_object(response: ClientResponse) -> dict[str, Any] | None: + try: + decoded = json.loads(response.body) + except (ValueError, UnicodeDecodeError): + return None + return decoded if isinstance(decoded, dict) else None __all__ = ["Recorder"] diff --git a/plugins/abridge/agentix/bridge/serve.py b/plugins/abridge/agentix/bridge/serve.py index dd44a90..ab1bf63 100644 --- a/plugins/abridge/agentix/bridge/serve.py +++ b/plugins/abridge/agentix/bridge/serve.py @@ -58,9 +58,12 @@ default keeps them for harvest. `--record-dir` (either mode) wraps each session's client in a `Recorder` -writing message-level rows to `/.jsonl`; rows +writing `abridge.record.v1` rows to `/.jsonl`; rows carry `session_id` + `request_id`, and the same `request_id` is stamped as `x-request-id` on the upstream hop so message rows join token records. +`--capture-level` picks how much of each call is persisted — `metadata` +(default: no conversation text) or `verbatim` (complete request body and +structured response). See `capture.py` for the schema. `GET /_health` reports `translation_spec_sha` — the SHA-256 of the Anthropic<->OpenAI transform module — so downstream data contracts can @@ -76,7 +79,6 @@ import argparse import asyncio import hashlib -import json import logging import os from collections import OrderedDict @@ -91,6 +93,7 @@ from fastapi import Request as FastAPIRequest from fastapi.responses import JSONResponse, Response +from .capture import CaptureLevel, canonical_text from .proxy import ( AbridgeError, Client, @@ -437,10 +440,23 @@ def _build_parser() -> argparse.ArgumentParser: "--record-dir", default=os.environ.get("ABRIDGE_RECORD_DIR"), help=( - "record every served (request, response) pair to " - "/.jsonl via Recorder — message-level rows with " - "session_id + request_id (+ gateway_session_id in tito mode), flushed " - "per line (env: ABRIDGE_RECORD_DIR)" + "record every served call to /.jsonl via Recorder " + "as abridge.record.v1 rows — session_id + request_id (+ gateway_session_id " + "in tito mode) and the prefix relation, flushed per line " + "(env: ABRIDGE_RECORD_DIR)" + ), + ) + parser.add_argument( + "--capture-level", + choices=[level.value for level in CaptureLevel], + default=os.environ.get("ABRIDGE_CAPTURE_LEVEL", CaptureLevel.METADATA.value), + help=( + "how much of each call --record-dir persists: 'metadata' (default) writes " + "identity, sampling, usage, tool names, per-message digests and the prefix " + "relation but no conversation text; 'verbatim' adds the complete request " + "body (system, full tool schemas, whole history) and the complete " + "structured response (thinking signatures, text, tool_use); 'off' records " + "nothing (env: ABRIDGE_CAPTURE_LEVEL)" ), ) parser.add_argument( @@ -455,28 +471,6 @@ def _build_parser() -> argparse.ArgumentParser: return parser -def _canonical_text(value: Any) -> str: - """Flatten Anthropic content (str or block list) to conversation-identity - text: text blocks contribute their text (cache_control and other - tokenization-irrelevant decorations are ignored), other blocks their - sorted-JSON form.""" - if value is None: - return "" - if isinstance(value, str): - return value - if isinstance(value, list): - parts: list[str] = [] - for block in value: - if isinstance(block, str): - parts.append(block) - elif isinstance(block, dict) and block.get("type") == "text": - parts.append(str(block.get("text", ""))) - else: - parts.append(json.dumps(block, sort_keys=True, ensure_ascii=False, default=repr)) - return "\n".join(parts) - return json.dumps(value, sort_keys=True, ensure_ascii=False, default=repr) - - def _conversation_key(body: dict[str, Any]) -> str: """A logical-conversation key for an Anthropic Messages request: the canonicalized system prompt + first user message. Turns of one @@ -487,7 +481,7 @@ def _conversation_key(body: dict[str, Any]) -> str: (m.get("content") for m in body.get("messages") or [] if isinstance(m, dict) and m.get("role") == "user"), None, ) - blob = _canonical_text(body.get("system")) + "\x00" + _canonical_text(first_user) + blob = canonical_text(body.get("system")) + "\x00" + canonical_text(first_user) return hashlib.sha256(blob.encode()).hexdigest()[:16] @@ -600,7 +594,10 @@ def build(session_id: str) -> Client: session_id=session_id, ) - if not args.record_dir: + level = CaptureLevel(getattr(args, "capture_level", CaptureLevel.METADATA.value)) + if not args.record_dir or level is CaptureLevel.OFF: + if args.record_dir: + logger.warning("abridge serve: --record-dir set with --capture-level off — nothing will be recorded") return build from .recorder import Recorder @@ -608,7 +605,9 @@ def build(session_id: str) -> Client: record_dir = Path(args.record_dir) def build_recorded(session_id: str) -> Client: - return Recorder(build(session_id), record_dir / f"{session_id}.jsonl", session_id=session_id) + return Recorder( + build(session_id), record_dir / f"{session_id}.jsonl", session_id=session_id, level=level + ) return build_recorded diff --git a/plugins/abridge/tests/test_anthropic_native_capture.py b/plugins/abridge/tests/test_anthropic_native_capture.py new file mode 100644 index 0000000..91fdeea --- /dev/null +++ b/plugins/abridge/tests/test_anthropic_native_capture.py @@ -0,0 +1,165 @@ +"""Capture on the Anthropic-native lane — the one a CLI coding agent uses. + +`AnthropicClient` forwards Messages verbatim to an Anthropic-protocol +endpoint and the agent asks for streaming, so the wire the agent gets back is +an SSE blob. The client publishes the completed `Message` it already holds, +which is what makes a `Recorder` row carry `thinking` signatures and +`tool_use` inputs as structure instead of a string somebody has to re-parse. +""" + +from __future__ import annotations + +import json + +import pytest +from agentix.bridge import CaptureLevel, Recorder, Request +from agentix.bridge.clients import AnthropicClient + +_SIGNATURE = "EqoBCkYIBRgCKkBz" * 12 + +_CONTENT = [ + {"type": "thinking", "thinking": "", "signature": _SIGNATURE}, + {"type": "text", "text": "listing the tree"}, + {"type": "tool_use", "id": "toolu_1", "name": "Bash", "input": {"command": "ls -la"}}, +] + +_BODY = { + "model": "claude-sonnet-4-5", + "max_tokens": 8192, + "stream": True, + "system": [{"type": "text", "text": "You are a coding assistant."}], + "tools": [ + { + "name": "Bash", + "description": "Run a shell command", + "input_schema": {"type": "object", "properties": {"command": {"type": "string"}}}, + } + ], + "messages": [{"role": "user", "content": "list the repo"}], +} + + +def _sdk_message(): + from anthropic.types import Message + + return Message.model_validate( + { + "id": "msg_1", + "type": "message", + "role": "assistant", + "model": "claude-sonnet-4-5", + "content": _CONTENT, + "stop_reason": "tool_use", + "stop_sequence": None, + "usage": {"input_tokens": 30112, "output_tokens": 214}, + } + ) + + +class _FakeStream: + """The shape `AsyncMessages.stream(...)` returns: an async CM that is also + an async iterator and can hand back the assembled final message.""" + + def __init__(self, message, calls: list[dict]) -> None: + self._message = message + self._calls = calls + + async def __aenter__(self) -> _FakeStream: + return self + + async def __aexit__(self, *_exc: object) -> bool: + return False + + def __aiter__(self) -> _FakeStream: + return self + + async def __anext__(self): + raise StopAsyncIteration + + async def get_final_message(self): + return self._message + + +def _recorded(tmp_path, *, level: CaptureLevel): + client = AnthropicClient(base_url="http://upstream.invalid", api_key="sk-ant-real", session_id="sess-1") + calls: list[dict] = [] + + def fake_stream(**kwargs): + calls.append(kwargs) + return _FakeStream(_sdk_message(), calls) + + client._client.messages.stream = fake_stream # noqa: SLF001 - stand in for the SDK hop + recorder = Recorder(client, tmp_path / "run.jsonl", session_id="sess-1", level=level) + return recorder.abridge_routes()["/v1/messages"], calls + + +def _rows(tmp_path) -> list[dict]: + return [json.loads(line) for line in (tmp_path / "run.jsonl").read_text().splitlines()] + + +@pytest.mark.asyncio +async def test_streaming_native_call_is_recorded_as_structure(tmp_path) -> None: + """The agent gets `text/event-stream`; the row gets the Message — with + the opaque `signature` preserved byte-for-byte, since it is the only + handle on the thinking content.""" + handler, _ = _recorded(tmp_path, level=CaptureLevel.VERBATIM) + response = await handler(Request(path="/v1/messages", body=_BODY)) + assert response.media_type == "text/event-stream" + + (row,) = _rows(tmp_path) + assert row["media_type"] == "text/event-stream" + assert "response_body" not in row + assert [b["type"] for b in row["response"]["content"]] == ["thinking", "text", "tool_use"] + assert row["response"]["content"][0]["signature"] == _SIGNATURE + assert row["response"]["content"][2]["input"] == {"command": "ls -la"} + assert row["stop_reason"] == "tool_use" + assert row["usage"]["input_tokens"] == 30112 + + +@pytest.mark.asyncio +async def test_streaming_native_call_records_the_full_request_body(tmp_path) -> None: + """`system`, the full tool schemas, and the history exactly as sent — + the ground truth everything downstream is derived from.""" + handler, upstream_calls = _recorded(tmp_path, level=CaptureLevel.VERBATIM) + await handler(Request(path="/v1/messages", body=_BODY)) + + (row,) = _rows(tmp_path) + assert row["request"] == _BODY + assert row["request"]["tools"][0]["input_schema"]["properties"] == {"command": {"type": "string"}} + assert row["shape"]["tool_names"] == ["Bash"] + # `stream` was popped off the body the SDK saw; the record keeps what the + # agent actually asked for. + assert "stream" not in upstream_calls[0] + assert row["stream"] is True + + +@pytest.mark.asyncio +async def test_native_client_reuses_the_recorders_request_id(tmp_path) -> None: + """The join key: the id in the row IS the `x-request-id` this client + stamps upstream, so a message-level row and a token-level record of the + same call line up.""" + handler, upstream_calls = _recorded(tmp_path, level=CaptureLevel.METADATA) + await handler(Request(path="/v1/messages", body=_BODY)) + + (row,) = _rows(tmp_path) + headers = upstream_calls[0]["extra_headers"] + assert headers["x-request-id"] == row["request_id"] + assert headers["x-session-id"] == "sess-1" + + +@pytest.mark.asyncio +async def test_metadata_level_keeps_the_signature_off_disk(tmp_path) -> None: + handler, _ = _recorded(tmp_path, level=CaptureLevel.METADATA) + await handler(Request(path="/v1/messages", body=_BODY)) + + (row,) = _rows(tmp_path) + blob = json.dumps(row) + assert _SIGNATURE not in blob and "ls -la" not in blob and "coding assistant" not in blob + # The shape still says a thinking block with an empty body and a long + # signature came back, and which tool the model called. + assert row["shape"]["content_blocks"][0] == { + "type": "thinking", + "chars": 0, + "signature_chars": len(_SIGNATURE), + } + assert row["shape"]["content_blocks"][2]["name"] == "Bash" diff --git a/plugins/abridge/tests/test_capture.py b/plugins/abridge/tests/test_capture.py new file mode 100644 index 0000000..d89ca73 --- /dev/null +++ b/plugins/abridge/tests/test_capture.py @@ -0,0 +1,256 @@ +"""`agentix.bridge.capture` — the derivations behind `abridge.record.v1`. + +These are the facts a downstream consumer reads instead of scraping a +harness's own logs: what tools the model actually saw, what the model +actually emitted, and whether the context was rewritten between two calls. +""" + +from __future__ import annotations + +import pytest +from agentix.bridge.capture import ( + CaptureLevel, + PrefixTracker, + canonical_text, + common_prefix_length, + digest, + request_facts, + response_facts, +) + + +def _tool(name: str, *, required: list[str] | None = None) -> dict: + return { + "name": name, + "description": f"the {name} tool", + "input_schema": { + "type": "object", + "properties": {"command": {"type": "string"}}, + "required": required or [], + }, + } + + +def _body(messages: list[dict], *, system: str = "be brief", tools: list[dict] | None = None) -> dict: + return { + "model": "claude-sonnet-4-5", + "max_tokens": 512, + "temperature": 1.0, + "system": system, + "tools": tools if tools is not None else [_tool("Bash"), _tool("Agent")], + "messages": messages, + } + + +# ── request facts ──────────────────────────────────────────────────────── + + +def test_tool_names_come_from_the_wire_body() -> None: + """The names recorded are the ones in the request the model saw — not + whatever a harness reports about its own tool surface elsewhere.""" + facts = request_facts(_body([{"role": "user", "content": "hi"}])) + assert facts.tool_names == ("Bash", "Agent") + assert facts.model == "claude-sonnet-4-5" + assert facts.stream is False + assert facts.sampling == {"max_tokens": 512, "temperature": 1.0} + + +def test_tools_digest_covers_the_schemas_not_just_the_names() -> None: + """A tool set can change without any name changing (a parameter becomes + required, a description is rewritten). `tool_names` alone would call that + identical; the digest does not.""" + same_names = request_facts(_body([], tools=[_tool("Bash"), _tool("Agent")])) + changed = request_facts(_body([], tools=[_tool("Bash", required=["command"]), _tool("Agent")])) + assert same_names.tool_names == changed.tool_names + assert same_names.tools_digest != changed.tools_digest + + +def test_conversation_key_ignores_cache_control_decorations() -> None: + """`cache_control` markers move between turns without changing what the + model reads, so they must not split a conversation into two lanes.""" + plain = request_facts(_body([], system=[{"type": "text", "text": "be brief"}])) + cached = request_facts( + _body([], system=[{"type": "text", "text": "be brief", "cache_control": {"type": "ephemeral"}}]) + ) + assert plain.conversation_key == cached.conversation_key + # The verbatim system prompt still differs, and the row says so. + assert plain.system_digest != cached.system_digest + + +def test_absent_system_and_tools_are_null_not_empty() -> None: + facts = request_facts({"model": "m", "messages": []}) + assert facts.system_digest is None + assert facts.tools_digest is None + assert facts.tool_names == () + assert facts.message_digests == () + + +def test_canonical_text_flattens_blocks_and_digest_is_key_order_stable() -> None: + assert canonical_text([{"type": "text", "text": "a"}, {"type": "text", "text": "b"}]) == "a\nb" + assert canonical_text(None) == "" and canonical_text("x") == "x" + assert digest({"a": 1, "b": 2}) == digest({"b": 2, "a": 1}) + + +# ── response facts ─────────────────────────────────────────────────────── + + +def test_thinking_block_shape_separates_text_from_signature() -> None: + """On this provider a thinking block arrives with empty text and a + populated opaque signature. `chars: 0, signature_chars: N` is the shape + that makes that visible at the metadata level.""" + facts = response_facts( + { + "content": [ + {"type": "thinking", "thinking": "", "signature": "x" * 210}, + {"type": "redacted_thinking", "data": "y" * 12}, + {"type": "text", "text": "hello"}, + {"type": "tool_use", "id": "toolu_1", "name": "Bash", "input": {"command": "ls"}}, + ], + "stop_reason": "tool_use", + "usage": {"input_tokens": 10, "output_tokens": 2}, + } + ) + assert facts.content_blocks == ( + {"type": "thinking", "chars": 0, "signature_chars": 210}, + {"type": "redacted_thinking", "data_chars": 12}, + {"type": "text", "chars": 5}, + {"type": "tool_use", "name": "Bash", "input_chars": len('{"command": "ls"}')}, + ) + assert facts.stop_reason == "tool_use" + assert facts.usage == {"input_tokens": 10, "output_tokens": 2} + + +def test_missing_message_yields_empty_facts() -> None: + facts = response_facts(None) + assert facts.content_blocks is None + assert facts.stop_reason is None and facts.usage is None + assert facts.response_digest is None and facts.assistant_digest is None + + +# ── prefix relation ────────────────────────────────────────────────────── + + +def _turn(tracker: PrefixTracker, request_id: str, messages: list[dict], *, system: str = "be brief"): + facts = request_facts(_body(messages, system=system)) + return facts, tracker.begin(request_id=request_id, facts=facts) + + +def test_common_prefix_length() -> None: + assert common_prefix_length(["a", "b", "c"], ["a", "b", "z"]) == 2 + assert common_prefix_length([], ["a"]) == 0 + assert common_prefix_length(["a"], ["a", "b"]) == 1 + + +def test_first_turn_has_nothing_to_extend() -> None: + _, relation = _turn(PrefixTracker(), "r1", [{"role": "user", "content": "one"}]) + assert relation.stable is True + assert relation.divergence_index is None + assert relation.previous_request_id is None + assert relation.assistant_echo is None + + +def test_an_appended_turn_is_stable() -> None: + tracker = PrefixTracker() + history = [{"role": "user", "content": "one"}] + _turn(tracker, "r1", history) + _, relation = _turn(tracker, "r2", [*history, {"role": "user", "content": "two"}]) + assert relation.stable is True + assert relation.divergence_index is None + assert relation.common_prefix == 1 + assert relation.previous_request_id == "r1" + assert relation.previous_messages == 1 + + +def test_a_rewritten_history_is_not_an_extension() -> None: + """What compaction looks like from here: same system prompt, a history + that no longer starts with what the previous request started with. No + boundary record, no `parentUuid`, no summary string — one boolean.""" + tracker = PrefixTracker() + _turn(tracker, "r1", [{"role": "user", "content": f"turn {i}"} for i in range(6)]) + _, relation = _turn(tracker, "r2", [{"role": "user", "content": ""}, + {"role": "user", "content": "turn 6"}]) + assert relation.stable is False + assert relation.divergence_index == 0 + assert relation.previous_messages == 6 + + +def test_a_partially_rewritten_history_reports_where_it_diverged() -> None: + tracker = PrefixTracker() + head = [{"role": "user", "content": "a"}, {"role": "user", "content": "b"}] + _turn(tracker, "r1", [*head, {"role": "user", "content": "c"}]) + _, relation = _turn(tracker, "r2", [*head, {"role": "user", "content": "REWRITTEN"}]) + assert relation.stable is False + assert relation.divergence_index == 2 + assert relation.common_prefix == 2 + + +def test_a_truncated_history_is_a_rewrite_too() -> None: + """Dropping the tail agrees as far as it goes but is not an extension — + splicing it into one stream would silently lose turns.""" + tracker = PrefixTracker() + history = [{"role": "user", "content": c} for c in "abc"] + _turn(tracker, "r1", history) + _, relation = _turn(tracker, "r2", history[:2]) + assert relation.stable is False + assert relation.divergence_index == 2 + + +def test_tool_and_system_changes_are_reported_inside_a_lane() -> None: + tracker = PrefixTracker() + history = [{"role": "user", "content": "one"}] + facts = request_facts(_body(history)) + tracker.begin(request_id="r1", facts=facts) + grown = request_facts(_body([*history, {"role": "user", "content": "two"}], tools=[_tool("Bash")])) + relation = tracker.begin(request_id="r2", facts=grown) + assert relation.stable is True + assert relation.tools_changed is True + assert relation.system_changed is False + + +def test_lanes_keep_interleaved_conversations_apart() -> None: + tracker = PrefixTracker() + main = [{"role": "user", "content": "main"}] + _turn(tracker, "r1", main, system="main agent") + _turn(tracker, "r2", [{"role": "user", "content": "sub"}], system="sub agent") + _, relation = _turn(tracker, "r3", [*main, {"role": "user", "content": "more"}], system="main agent") + assert relation.stable is True + assert relation.previous_request_id == "r1" + + +def test_lane_eviction_reads_as_a_fresh_conversation() -> None: + """A bounded tracker forgets the least recent lane. The next request on + it looks like a first turn — the safe direction, since the alternative is + comparing against a lane that no longer exists.""" + tracker = PrefixTracker(max_lanes=2) + _turn(tracker, "r1", [{"role": "user", "content": "a"}], system="one") + _turn(tracker, "r2", [{"role": "user", "content": "b"}], system="two") + _turn(tracker, "r3", [{"role": "user", "content": "c"}], system="three") + _, relation = _turn(tracker, "r4", [{"role": "user", "content": "a"}], system="one") + assert relation.previous_request_id is None + + +def test_finish_ignores_a_lane_another_call_already_took_over() -> None: + """Two calls racing on one lane must not let the loser's response be + attributed to the winner's baseline.""" + tracker = PrefixTracker() + history = [{"role": "user", "content": "one"}] + facts, _ = _turn(tracker, "r1", history) + _turn(tracker, "r2", [*history, {"role": "user", "content": "two"}]) + tracker.finish(request_id="r1", conversation_key=facts.conversation_key, assistant_digest="stale") + + _, relation = _turn(tracker, "r3", [*history, {"role": "user", "content": "two"}, + {"role": "user", "content": "three"}]) + assert relation.assistant_echo is None # r2 never reported one; r1's is not borrowed + + +def test_max_lanes_must_be_positive() -> None: + with pytest.raises(ValueError, match="max_lanes"): + PrefixTracker(max_lanes=0) + + +# ── levels ─────────────────────────────────────────────────────────────── + + +def test_capture_levels_are_a_named_ladder() -> None: + assert [level.value for level in CaptureLevel] == ["off", "metadata", "verbatim"] + assert CaptureLevel("verbatim") is CaptureLevel.VERBATIM diff --git a/plugins/abridge/tests/test_recorder.py b/plugins/abridge/tests/test_recorder.py index 395c475..9894a50 100644 --- a/plugins/abridge/tests/test_recorder.py +++ b/plugins/abridge/tests/test_recorder.py @@ -1,29 +1,76 @@ """`Recorder` — host-side rollout capture at the tunnel. -Wrapping a client records every (request, response) pair its handlers -serve to a JSONL file, without the agent or the upstream noticing. The -tunnel is the one place all of an agent's LLM traffic passes, so this is +Wrapping a client records every call its handlers serve as an +`abridge.record.v1` JSONL row, without the agent or the upstream noticing. +The tunnel is the one place all of an agent's LLM traffic passes, so this is the natural recording point for rollout data collection. + +Two properties carry the weight here: the level ladder (metadata never puts +conversation text on disk; verbatim is the reconstruction ground truth) and +the prefix relation (a row says whether the context was rewritten, so nobody +downstream has to parse a harness's own boundary records). """ from __future__ import annotations import json +import stat from pathlib import Path import pytest -from agentix.bridge import AbridgeError, ClientResponse, Recorder, Request, on +from agentix.bridge import AbridgeError, CaptureLevel, ClientResponse, Recorder, Request, on +from agentix.bridge._request_id import publish_response_message + +MESSAGES = "/v1/messages" +COUNT_TOKENS = "/v1/messages/count_tokens" + +_TOOLS = [ + { + "name": "Bash", + "description": "Run a shell command", + "input_schema": {"type": "object", "properties": {"command": {"type": "string"}}}, + } +] + + +def _body(*, messages: list[dict], system: str = "be brief", tools: list[dict] | None = None) -> dict: + return { + "model": "claude-sonnet-4-5", + "max_tokens": 512, + "stream": True, + "system": system, + "tools": _TOOLS if tools is None else tools, + "messages": messages, + } + + +def _message(content: list[dict], *, stop_reason: str = "end_turn") -> dict: + return { + "id": "msg_1", + "type": "message", + "role": "assistant", + "model": "claude-sonnet-4-5", + "content": content, + "stop_reason": stop_reason, + "usage": {"input_tokens": 30112, "output_tokens": 214}, + } class _EchoClient: - def __init__(self) -> None: + """An Anthropic-face client that streams (like the real CLI does) and + publishes the completed Message the way the bundled clients do.""" + + def __init__(self, content: list[dict] | None = None) -> None: self.closed = False + self.content = content if content is not None else [{"type": "text", "text": "hi back"}] - @on("/v1/messages") + @on(MESSAGES) async def messages(self, request: Request) -> ClientResponse: - return ClientResponse.json({"echo": request.body["msg"], "id": "resp-1"}) + message = _message(self.content) + publish_response_message(message) + return ClientResponse.sse(b"event: message_stop\ndata: {}\n\n") - @on("/v1/messages/count_tokens") + @on(COUNT_TOKENS) async def count_tokens(self, request: Request) -> ClientResponse: return ClientResponse.json({"input_tokens": 7}) @@ -32,120 +79,292 @@ async def aclose(self) -> None: class _FailingClient: - @on("/v1/messages") + @on(MESSAGES) async def messages(self, request: Request) -> ClientResponse: raise AbridgeError("upstream exploded", status_code=502) -def _lines(path: Path) -> list[dict]: +def _rows(path: Path) -> list[dict]: return [json.loads(line) for line in path.read_text().splitlines()] -@pytest.mark.asyncio -async def test_recorder_exposes_inner_routes_and_records_pairs(tmp_path) -> None: - out = tmp_path / "run.jsonl" - recorder = Recorder(_EchoClient(), out) - routes = recorder.abridge_routes() - assert set(routes) == {"/v1/messages", "/v1/messages/count_tokens"} +def _turns(path: Path) -> list[dict]: + return [r for r in _rows(path) if r["schema_version"] == "abridge.record.v1"] - resp = await routes["/v1/messages"](Request(path="/v1/messages", body={"msg": "hi"})) - assert json.loads(resp.body) == {"echo": "hi", "id": "resp-1"} - (record,) = _lines(out) - assert record["path"] == "/v1/messages" - assert record["request"] == {"msg": "hi"} - assert record["response"]["body"] == {"echo": "hi", "id": "resp-1"} - assert record["response"]["status_code"] == 200 - assert "ts" in record +async def _call(routes, body: dict, path: str = MESSAGES) -> ClientResponse: + return await routes[path](Request(path=path, body=body)) + + +# ── levels ─────────────────────────────────────────────────────────────── @pytest.mark.asyncio -async def test_recorder_appends_one_line_per_call_in_order(tmp_path) -> None: +async def test_metadata_level_is_the_default_and_writes_no_conversation_text(tmp_path) -> None: + """Turning capture on must not turn verbatim capture on: the default + level keeps the join keys, the shape, and the prefix relation, and puts + no prompt or completion text on disk.""" out = tmp_path / "run.jsonl" recorder = Recorder(_EchoClient(), out) + assert recorder.level is CaptureLevel.METADATA routes = recorder.abridge_routes() - for i in range(3): - await routes["/v1/messages"](Request(path="/v1/messages", body={"msg": f"m{i}"})) - assert [r["request"]["msg"] for r in _lines(out)] == ["m0", "m1", "m2"] + assert set(routes) == {MESSAGES, COUNT_TOKENS} + + await _call(routes, _body(messages=[{"role": "user", "content": "the secret plan"}])) + + (row,) = _turns(out) + blob = json.dumps(row) + assert row["capture_level"] == "metadata" + assert "request" not in row and "response" not in row and "response_body" not in row + assert "the secret plan" not in blob and "be brief" not in blob and "hi back" not in blob + # The parts that are NOT conversation text still land. + assert row["model"] == "claude-sonnet-4-5" + assert row["stream"] is True + assert row["sampling"] == {"max_tokens": 512} + assert row["shape"]["tool_names"] == ["Bash"] + assert row["shape"]["messages"] == 1 + assert row["shape"]["content_blocks"] == [{"type": "text", "chars": 7}] + assert row["stop_reason"] == "end_turn" + assert row["usage"] == {"input_tokens": 30112, "output_tokens": 214} @pytest.mark.asyncio -async def test_recorder_records_handler_errors_and_reraises(tmp_path) -> None: +async def test_verbatim_level_records_the_complete_request_body(tmp_path) -> None: + """The reconstruction ground truth: system, the FULL tool schemas (not + just their names), and the entire message history exactly as sent.""" out = tmp_path / "run.jsonl" - recorder = Recorder(_FailingClient(), out) - routes = recorder.abridge_routes() - with pytest.raises(AbridgeError): - await routes["/v1/messages"](Request(path="/v1/messages", body={"msg": "boom"})) + routes = Recorder(_EchoClient(), out, level=CaptureLevel.VERBATIM).abridge_routes() + body = _body(messages=[{"role": "user", "content": [{"type": "text", "text": "hello"}]}]) + + await _call(routes, body) - (record,) = _lines(out) - assert record["request"] == {"msg": "boom"} - assert "upstream exploded" in record["error"] - assert "response" not in record + (row,) = _turns(out) + assert row["capture_level"] == "verbatim" + assert row["request"] == body + assert row["request"]["tools"][0]["input_schema"]["properties"] == {"command": {"type": "string"}} @pytest.mark.asyncio -async def test_recorder_aclose_closes_inner_and_flushes(tmp_path) -> None: +async def test_verbatim_level_records_thinking_signatures_verbatim(tmp_path) -> None: + """This provider returns `thinking` blocks with empty text and a + populated opaque `signature` — the only handle on that content. It has to + survive byte-for-byte, and the skeleton has to show the shape.""" out = tmp_path / "run.jsonl" - inner = _EchoClient() - recorder = Recorder(inner, out) - routes = recorder.abridge_routes() - await routes["/v1/messages"](Request(path="/v1/messages", body={"msg": "x"})) - await recorder.aclose() - assert inner.closed - assert len(_lines(out)) == 1 + content = [ + {"type": "thinking", "thinking": "", "signature": "EqoBCkYIBRgCKkC" * 8}, + {"type": "text", "text": "running it"}, + {"type": "tool_use", "id": "toolu_1", "name": "Bash", "input": {"command": "ls"}}, + ] + routes = Recorder(_EchoClient(content), out, level=CaptureLevel.VERBATIM).abridge_routes() + + await _call(routes, _body(messages=[{"role": "user", "content": "go"}])) + + (row,) = _turns(out) + assert row["response"]["content"] == content + assert row["shape"]["content_blocks"] == [ + {"type": "thinking", "chars": 0, "signature_chars": len(content[0]["signature"])}, + {"type": "text", "chars": 10}, + {"type": "tool_use", "name": "Bash", "input_chars": len(json.dumps({"command": "ls"}))}, + ] + + +@pytest.mark.asyncio +async def test_streaming_response_is_recorded_as_structure_not_an_sse_blob(tmp_path) -> None: + """The agent's wire is `text/event-stream`, but the row carries the + Message object the client published — nobody downstream should have to + re-parse an SSE string to find a tool_use block.""" + out = tmp_path / "run.jsonl" + routes = Recorder(_EchoClient(), out, level=CaptureLevel.VERBATIM).abridge_routes() + + await _call(routes, _body(messages=[{"role": "user", "content": "go"}])) + + (row,) = _turns(out) + assert row["media_type"] == "text/event-stream" + assert row["response"]["content"] == [{"type": "text", "text": "hi back"}] + assert "response_body" not in row @pytest.mark.asyncio -async def test_recorder_preserves_non_json_bodies_as_text(tmp_path) -> None: - class _SseClient: - @on("/v1/messages") +async def test_opaque_body_is_kept_as_text_and_flagged_by_the_missing_structure(tmp_path) -> None: + """A bare pass-through (nobody published a Message) can only be recorded + as bytes. That boundary is stated — `response` absent, `content_blocks` + null — instead of being papered over with a guess.""" + + class _Relay: + @on(MESSAGES) async def messages(self, request: Request) -> ClientResponse: return ClientResponse.sse(b"event: ping\ndata: {}\n\n") out = tmp_path / "run.jsonl" - routes = Recorder(_SseClient(), out).abridge_routes() - await routes["/v1/messages"](Request(path="/v1/messages", body={})) - (record,) = _lines(out) - assert record["response"]["media_type"] == "text/event-stream" - assert record["response"]["body"] == "event: ping\ndata: {}\n\n" + routes = Recorder(_Relay(), out, level=CaptureLevel.VERBATIM).abridge_routes() + await _call(routes, _body(messages=[{"role": "user", "content": "go"}])) + (row,) = _turns(out) + assert "response" not in row + assert row["response_body"] == "event: ping\ndata: {}\n\n" + assert row["shape"]["content_blocks"] is None + assert row["stop_reason"] is None -def test_recorder_delegates_environ(tmp_path) -> None: - class _EnvClient(_EchoClient): - def environ(self, handle) -> dict[str, str]: - return {"X": "y"} - recorder = Recorder(_EnvClient(), tmp_path / "run.jsonl") - assert recorder.environ(None) == {"X": "y"} +@pytest.mark.asyncio +async def test_count_tokens_stays_at_metadata_even_at_verbatim(tmp_path) -> None: + """A token-count call re-sends a duplicate of the adjacent history and + produces no completion. Its body is excluded from verbatim capture, and + the row says `metadata` rather than looking like a verbatim row that lost + its body.""" + out = tmp_path / "run.jsonl" + routes = Recorder(_EchoClient(), out, level=CaptureLevel.VERBATIM).abridge_routes() + + await _call(routes, _body(messages=[{"role": "user", "content": "count me"}]), COUNT_TOKENS) + + (row,) = _turns(out) + assert row["path"] == COUNT_TOKENS + assert row["capture_level"] == "metadata" + assert "request" not in row and "count me" not in json.dumps(row) + + +@pytest.mark.asyncio +async def test_non_anthropic_paths_keep_their_verbatim_body_without_derivations(tmp_path) -> None: + """`Recorder` is not Anthropic-only: another face's body is still + captured verbatim, it just carries no Anthropic-face derivations.""" + + class _OpenAIish: + @on("/v1/chat/completions") + async def chat(self, request: Request) -> ClientResponse: + return ClientResponse.json({"choices": [{"message": {"content": "ok"}}]}) + + out = tmp_path / "run.jsonl" + routes = Recorder(_OpenAIish(), out, level=CaptureLevel.VERBATIM).abridge_routes() + await routes["/v1/chat/completions"]( + Request(path="/v1/chat/completions", body={"model": "m", "messages": [{"role": "user", "content": "hi"}]}) + ) + + (row,) = _turns(out) + assert row["request"]["messages"] == [{"role": "user", "content": "hi"}] + assert row["response"] == {"choices": [{"message": {"content": "ok"}}]} + assert "shape" not in row and "prefix" not in row + + +def test_off_level_is_rejected_rather_than_writing_an_empty_file(tmp_path) -> None: + """`off` means no recorder at all — building one that writes nothing is a + configuration mistake worth failing loudly on.""" + with pytest.raises(ValueError, match="CaptureLevel.OFF"): + Recorder(_EchoClient(), tmp_path / "run.jsonl", level=CaptureLevel.OFF) + + +# ── prefix relation ────────────────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_appended_turns_are_stable_and_a_rewrite_is_not(tmp_path) -> None: + """The load-bearing fact. Turn 2 extends turn 1 -> stable. Turn 3 + replaces the history (what compaction does) -> `stable: false` with the + index where the two diverge, and no boundary record had to be parsed.""" + out = tmp_path / "run.jsonl" + routes = Recorder(_EchoClient(), out).abridge_routes() + + first = [{"role": "user", "content": "one"}] + await _call(routes, _body(messages=first)) + grown = [*first, {"role": "assistant", "content": [{"type": "text", "text": "hi back"}]}, + {"role": "user", "content": "two"}] + await _call(routes, _body(messages=grown)) + compacted = [{"role": "user", "content": ""}, + {"role": "user", "content": "three"}] + await _call(routes, _body(messages=compacted)) + + one, two, three = _turns(out) + assert one["prefix"]["stable"] is True and one["prefix"]["previous_request_id"] is None + assert two["prefix"]["stable"] is True + assert two["prefix"]["previous_request_id"] == one["request_id"] + assert three["prefix"] == { + "stable": False, + "divergence_index": 0, + "common_prefix": 0, + "previous_request_id": two["request_id"], + "previous_messages": 3, + "system_changed": False, + "tools_changed": False, + "assistant_echo": None, + } + + +@pytest.mark.asyncio +async def test_assistant_echo_reports_whether_the_model_turn_came_back_intact(tmp_path) -> None: + """Whether the harness echoed the assistant turn we returned decides + whether the rows splice into one stream. `verbatim` when it came back + byte-identical, `modified` when the harness rewrote it.""" + out = tmp_path / "run.jsonl" + content = [{"type": "thinking", "thinking": "", "signature": "sig"}, {"type": "text", "text": "hi back"}] + routes = Recorder(_EchoClient(content), out).abridge_routes() + first = [{"role": "user", "content": "one"}] + + # Lane A echoes the assistant turn byte-for-byte; lane B drops the + # thinking block; lane C re-sends the same history without it at all. + for system in ("agent A", "agent B", "agent C"): + await _call(routes, _body(messages=first, system=system)) + await _call(routes, _body(system="agent A", messages=[*first, {"role": "assistant", "content": content}, + {"role": "user", "content": "two"}])) + stripped = [{"type": "text", "text": "hi back"}] + await _call(routes, _body(system="agent B", messages=[*first, {"role": "assistant", "content": stripped}, + {"role": "user", "content": "two"}])) + await _call(routes, _body(system="agent C", messages=first)) + + *_, echoed, altered, retried = _turns(out) + assert [r["prefix"]["stable"] for r in (echoed, altered, retried)] == [True, True, True] + assert echoed["prefix"]["assistant_echo"] == "verbatim" + assert altered["prefix"]["assistant_echo"] == "modified" + assert retried["prefix"]["assistant_echo"] == "absent" + + +@pytest.mark.asyncio +async def test_interleaved_conversations_do_not_report_false_rewrites(tmp_path) -> None: + """One key multiplexes several conversations (subagents, helper calls). + They are tracked in separate lanes keyed by the system prompt, so + alternating between them is not mistaken for a rewrite.""" + out = tmp_path / "run.jsonl" + routes = Recorder(_EchoClient(), out).abridge_routes() + + main = [{"role": "user", "content": "main one"}] + sub = [{"role": "user", "content": "sub one"}] + await _call(routes, _body(messages=main, system="main agent")) + await _call(routes, _body(messages=sub, system="sub agent")) + await _call(routes, _body(messages=[*main, {"role": "user", "content": "main two"}], system="main agent")) + + rows = _turns(out) + assert [r["prefix"]["stable"] for r in rows] == [True, True, True] + assert rows[0]["conversation_key"] == rows[2]["conversation_key"] != rows[1]["conversation_key"] + assert rows[2]["prefix"]["previous_request_id"] == rows[0]["request_id"] + + +# ── identity, ordering, durability ─────────────────────────────────────── @pytest.mark.asyncio -async def test_recorder_rows_carry_session_and_request_ids(tmp_path) -> None: +async def test_rows_carry_session_and_request_ids_and_a_monotonic_turn_index(tmp_path) -> None: """Rows are joinable against downstream token records: `session_id` (the rollout identity the Recorder was built with) and a per-call - `request_id`, unique across calls.""" + `request_id`, unique across calls; `turn_index` orders them.""" out = tmp_path / "run.jsonl" - recorder = Recorder(_EchoClient(), out, session_id="sess-42") - routes = recorder.abridge_routes() - await routes["/v1/messages"](Request(path="/v1/messages", body={"msg": "a"})) - await routes["/v1/messages"](Request(path="/v1/messages", body={"msg": "b"})) + routes = Recorder(_EchoClient(), out, session_id="sess-42").abridge_routes() + for i in range(3): + await _call(routes, _body(messages=[{"role": "user", "content": f"m{i}"}])) - first, second = _lines(out) - assert first["session_id"] == second["session_id"] == "sess-42" - assert first["request_id"] and second["request_id"] - assert first["request_id"] != second["request_id"] + rows = _turns(out) + assert {r["session_id"] for r in rows} == {"sess-42"} + assert [r["turn_index"] for r in rows] == [0, 1, 2] + assert len({r["request_id"] for r in rows}) == 3 @pytest.mark.asyncio -async def test_recorder_request_id_matches_upstream_x_request_id(tmp_path, monkeypatch) -> None: - """The alignment contract: the id in the Recorder row IS the - `x-request-id` the transport stamps on the upstream hop (via the - `current_request_id` context var), so a message-level row and a - token-level sidecar record for the same call join on one key.""" +async def test_request_id_matches_the_upstream_x_request_id(tmp_path, monkeypatch) -> None: + """The alignment contract: the id in the row IS the `x-request-id` the + transport stamps on the upstream hop (via the `current_request_id` + context var), so a message-level row and a token-level sidecar record for + the same call join on one key.""" import httpx from agentix.bridge import Forward - fwd = Forward("http://side.car", paths=["/v1/messages"], session_id="sess-1") + fwd = Forward("http://side.car", paths=[MESSAGES], session_id="sess-1") seen_headers: list[dict] = [] async def fake_post(url, *, json, headers): @@ -156,57 +375,115 @@ async def fake_post(url, *, json, headers): out = tmp_path / "run.jsonl" routes = Recorder(fwd, out, session_id="sess-1").abridge_routes() - await routes["/v1/messages"](Request(path="/v1/messages", body={"x": 1})) + await _call(routes, _body(messages=[{"role": "user", "content": "x"}])) - (row,) = _lines(out) + (row,) = _turns(out) assert seen_headers[0]["x-request-id"] == row["request_id"] assert seen_headers[0]["x-session-id"] == row["session_id"] == "sess-1" @pytest.mark.asyncio -async def test_recorder_error_rows_also_carry_ids(tmp_path) -> None: +async def test_handler_errors_are_recorded_and_reraised(tmp_path) -> None: out = tmp_path / "run.jsonl" - routes = Recorder(_FailingClient(), out, session_id="sess-9").abridge_routes() + routes = Recorder(_FailingClient(), out, session_id="sess-9", level=CaptureLevel.VERBATIM).abridge_routes() with pytest.raises(AbridgeError): - await routes["/v1/messages"](Request(path="/v1/messages", body={})) - (row,) = _lines(out) - assert row["session_id"] == "sess-9" - assert row["request_id"] - assert "error" in row + await _call(routes, _body(messages=[{"role": "user", "content": "boom"}])) + + (row,) = _turns(out) + assert "upstream exploded" in row["error"] + assert "response" not in row and "status_code" not in row + # The request side still landed — a failed call is signal, not a hole. + assert row["request"]["messages"] == [{"role": "user", "content": "boom"}] + assert row["session_id"] == "sess-9" and row["request_id"] + + +@pytest.mark.asyncio +async def test_aclose_closes_the_inner_client_and_seals_the_file(tmp_path) -> None: + """The trailer distinguishes a cleanly closed file from one truncated by + a process that died mid-rollout.""" + out = tmp_path / "run.jsonl" + inner = _EchoClient() + recorder = Recorder(inner, out, session_id="sess-1") + routes = recorder.abridge_routes() + await _call(routes, _body(messages=[{"role": "user", "content": "x"}])) + await recorder.aclose() + + assert inner.closed + *_, trailer = _rows(out) + assert trailer == { + "schema_version": "abridge.session.v1", + "session_id": "sess-1", + "turns": 1, + "capture_level": "metadata", + "ts": trailer["ts"], + } + + +def test_recorder_delegates_environ(tmp_path) -> None: + class _EnvClient(_EchoClient): + def environ(self, handle) -> dict[str, str]: + return {"X": "y"} + + recorder = Recorder(_EnvClient(), tmp_path / "run.jsonl") + assert recorder.environ(None) == {"X": "y"} -def test_recorder_opens_file_lazily(tmp_path) -> None: +@pytest.mark.asyncio +async def test_recorder_opens_file_lazily(tmp_path) -> None: """A Recorder that never serves (e.g. build_session_app's route - enumeration probe) must leave no empty file behind.""" + enumeration probe) must leave no empty file behind — not even a trailer.""" out = tmp_path / "probe.jsonl" recorder = Recorder(_EchoClient(), out) recorder.abridge_routes() + await recorder.aclose() assert not out.exists() @pytest.mark.asyncio -async def test_recorder_write_failure_is_log_and_serve(tmp_path, monkeypatch, caplog) -> None: - """Capture failures never fail the served call (matching the token - gateway's policy): with the record path unwritable the agent still gets - its response and the drop is logged.""" +async def test_record_file_is_created_private(tmp_path) -> None: + """A verbatim row holds whatever the user typed; the file must not be + readable by other accounts on the box.""" + out = tmp_path / "run.jsonl" + routes = Recorder(_EchoClient(), out, level=CaptureLevel.VERBATIM).abridge_routes() + await _call(routes, _body(messages=[{"role": "user", "content": "x"}])) + assert stat.S_IMODE(out.stat().st_mode) == 0o600 + + +@pytest.mark.asyncio +async def test_unserializable_row_is_log_and_serve_and_leaves_an_index_gap(tmp_path, caplog) -> None: + """Rows are strict JSON — no `NaN`/`Infinity` literals — so one parser + reads these and the token gateway's records. A row that violates it is + dropped, not silently emitted as non-standard JSON, and the served call + still succeeds: `turn_index` advances so the hole is detectable.""" import logging - out = tmp_path / "run.jsonl" - recorder = Recorder(_EchoClient(), out) - routes = recorder.abridge_routes() + class _NaNOnSecondCall: + calls = 0 + + @on(MESSAGES) + async def messages(self, request: Request) -> ClientResponse: + type(self).calls += 1 + message = _message([{"type": "text", "text": "hi"}]) + if type(self).calls == 2: + message["usage"] = {"input_tokens": float("nan")} + publish_response_message(message) + return ClientResponse.json(message) - def broken_open(*args, **kwargs): - raise OSError("disk full") + out = tmp_path / "run.jsonl" + routes = Recorder(_NaNOnSecondCall(), out).abridge_routes() - monkeypatch.setattr(type(out), "open", broken_open) + await _call(routes, _body(messages=[{"role": "user", "content": "one"}])) with caplog.at_level(logging.ERROR, logger="agentix.bridge.recorder"): - resp = await routes["/v1/messages"](Request(path="/v1/messages", body={"msg": "hi"})) - assert json.loads(resp.body)["echo"] == "hi" # the call succeeded + response = await _call(routes, _body(messages=[{"role": "user", "content": "two"}])) + assert response.status_code == 200 # the agent's call still succeeded assert any("row NOT persisted" in r.message for r in caplog.records) + await _call(routes, _body(messages=[{"role": "user", "content": "three"}])) + assert [r["turn_index"] for r in _turns(out)] == [0, 2] # index 1 is the detectable hole + @pytest.mark.asyncio -async def test_recorder_drops_rows_after_aclose(tmp_path, caplog) -> None: +async def test_rows_are_dropped_after_aclose(tmp_path, caplog) -> None: """A straggler dispatch that outlives aclose() must not resurrect the record file: the row is dropped (logged), the file stays closed, and no orphan handle is created.""" @@ -215,12 +492,12 @@ async def test_recorder_drops_rows_after_aclose(tmp_path, caplog) -> None: out = tmp_path / "run.jsonl" recorder = Recorder(_EchoClient(), out) routes = recorder.abridge_routes() - await routes["/v1/messages"](Request(path="/v1/messages", body={"msg": "before"})) + await _call(routes, _body(messages=[{"role": "user", "content": "before"}])) await recorder.aclose() with caplog.at_level(logging.WARNING, logger="agentix.bridge.recorder"): - resp = await routes["/v1/messages"](Request(path="/v1/messages", body={"msg": "late"})) - assert json.loads(resp.body)["echo"] == "late" # still served + response = await _call(routes, _body(messages=[{"role": "user", "content": "late"}])) + assert response.status_code == 200 # still served assert any("recorder is closed" in r.message for r in caplog.records) - assert [r["request"]["msg"] for r in _lines(out)] == ["before"] # no late row + assert [r["turn_index"] for r in _turns(out)] == [0] # no late row assert recorder._file is not None and recorder._file.closed # noqa: SLF001 - not reopened diff --git a/plugins/abridge/tests/test_tito_composition.py b/plugins/abridge/tests/test_tito_composition.py index ed1a971..7d94a26 100644 --- a/plugins/abridge/tests/test_tito_composition.py +++ b/plugins/abridge/tests/test_tito_composition.py @@ -212,7 +212,7 @@ def test_tito_composition_with_record_dir_joins_rows_to_gateway_calls(fake_tito: """--record-dir in tito mode: the message-level Recorder row and the gateway's x-request-id share one id, and the row's session_id is the caller-derived serve session (the row->record join key set).""" - tc = TestClient(_tito_app(fake_tito, "--record-dir", str(tmp_path))) + tc = TestClient(_tito_app(fake_tito, "--record-dir", str(tmp_path), "--capture-level", "verbatim")) r = tc.post("/v1/messages", json=_ANTHROPIC_BODY, headers={"x-api-key": "rollout-1"}) assert r.status_code == 200 @@ -227,11 +227,53 @@ def test_tito_composition_with_record_dir_joins_rows_to_gateway_calls(fake_tito: assert row["gateway_session_id"] == call["session"] == _FakeTito.sessions[0] assert row["path"] == "/v1/messages" assert row["request"] == _ANTHROPIC_BODY # the agent-side (Anthropic) shape - assert row["response"]["body"]["content"] == [{"type": "text", "text": "hello from tito"}] + assert row["response"]["content"] == [{"type": "text", "text": "hello from tito"}] # No file for the route-enumeration probe session. assert sorted(p.name for p in tmp_path.iterdir()) == [f"{serve_session}.jsonl"] +def test_tito_composition_record_dir_defaults_to_metadata_level(fake_tito: str, tmp_path) -> None: + """--record-dir alone never puts prompts on disk: the default level keeps + the join keys and the prefix relation and drops the bodies. Turning + capture on must not silently turn verbatim capture on.""" + tc = TestClient(_tito_app(fake_tito, "--record-dir", str(tmp_path))) + assert tc.post("/v1/messages", json=_ANTHROPIC_BODY, headers={"x-api-key": "rollout-1"}).status_code == 200 + + serve_session = session_id_for("rollout-1") + (row,) = [json.loads(line) for line in (tmp_path / f"{serve_session}.jsonl").read_text().splitlines()] + assert row["capture_level"] == "metadata" + assert "request" not in row and "response" not in row and "response_body" not in row + assert row["request_id"] and row["gateway_session_id"] == _FakeTito.sessions[0] + assert row["prefix"]["stable"] is True + assert "be brief" not in json.dumps(row) and "hello from tito" not in json.dumps(row) + + +def test_record_rows_never_contain_caller_credentials(fake_tito: str, tmp_path) -> None: + """Even at the verbatim level a record cannot leak a credential: the + tunnel carries no HTTP metadata at all and `serve` only ever hashes the + caller key into a session id. That is a structural guarantee, not a + redaction pass that could miss a field.""" + secret = "rollout-secret-do-not-log" + tc = TestClient(_tito_app(fake_tito, "--record-dir", str(tmp_path), "--capture-level", "verbatim")) + r = tc.post( + "/v1/messages", + json=_ANTHROPIC_BODY, + headers={"authorization": f"Bearer {secret}", "x-forwarded-for": "10.1.2.3"}, + ) + assert r.status_code == 200 + + text = (tmp_path / f"{session_id_for(secret)}.jsonl").read_text() + assert secret not in text + assert "authorization" not in text.lower() + assert "10.1.2.3" not in text + + +def test_capture_level_off_records_nothing(fake_tito: str, tmp_path) -> None: + tc = TestClient(_tito_app(fake_tito, "--record-dir", str(tmp_path), "--capture-level", "off")) + assert tc.post("/v1/messages", json=_ANTHROPIC_BODY, headers={"x-api-key": "r1"}).status_code == 200 + assert list(tmp_path.iterdir()) == [] + + def test_tito_streaming_agent_gets_replayed_sse(fake_tito: str) -> None: """stream:true agents get the locally rendered SSE replay while the gateway still saw a non-streaming call.""" diff --git a/plugins/tito/README.md b/plugins/tito/README.md index 7b261f3..76c5175 100644 --- a/plugins/tito/README.md +++ b/plugins/tito/README.md @@ -31,6 +31,43 @@ The algorithm is **model-agnostic** (base `TITOTokenizer`); a model family is a fixed chat template plus a tiny boundary fixup — e.g. `Qwen3TITOTokenizer` re-inserts the `\n` after `<|im_end|>` that the model omits when it stops. +Families (`--tito-model`): + +- `default` — the tokenizer's own template, no fixups. +- `qwen3` — bundled fixed Qwen3 template (`qwen3_fixed.jinja`) + newline fixup. +- `qwen3_5` — Qwen3.5 **and Qwen3.8** (`Qwen3_5ForConditionalGeneration`): + the tokenizer's own template + the Qwen3 newline fixup + a synthetic + context that carries a dummy user turn, because this template family + raises `No user query found in messages` otherwise. Tool calls are the + `` dialect (vLLM `--tool-call-parser qwen3_coder`), + reasoning comes from `reasoning_content` only. Pin the template variables + that change the prompt with `--tito-chat-template-kwargs`, e.g. + `'{"reasoning_effort": "xhigh"}'` for Qwen3.8 (its default is `xhigh`; + `medium`/`low` are the other legal values). Appending `system` messages + mid-conversation is rejected; `user` appends are safe on Qwen3.8 (whose + template preserves earlier thinking by default) but clear earlier-turn + reasoning on Qwen3.5, so keep `--tito-allowed-append-roles tool` there. + Three agent-facing tolerances learned from Pi 0.84.1 driving Qwen3.8: an + assistant turn that spent its whole `max_tokens` inside `` (reasoning, + no content, no tool calls) is a legal turn, not a malformed upstream reply; + echoed `tool_calls` are compared by their **parsed** arguments (Pi + re-serializes JSON with different spacing than vLLM); and a history that + echoes reasoning under `reasoning` (vLLM's key) is aliased to + `reasoning_content` before rendering so earlier-turn thinking survives the + from-scratch render audit. + +## Streaming clients + +The gateway forces `stream=false` upstream (the token harvest needs the full +JSON body). A client that sent `stream: true` — Pi's `openai-completions` +provider, the OpenAI SDK with `stream=True` — gets the **completed turn +re-framed as Server-Sent Events**: one chunk with the whole assistant delta +(`role`, `content`, the backend's reasoning field, `tool_calls` with `index`), +one terminal chunk with `finish_reason` and `usage`, then `data: [DONE]`. The +response carries `x-tito-stream: reframed`; content is identical to the JSON +body, only the framing changes. Non-200 upstream answers pass through as +JSON regardless of framing. + ## Backend kinds The token dialect is selected by `--backend-kind` (`TITOGatewayConfig.backend_kind`): @@ -72,12 +109,27 @@ agentix-tito serve \ --session-server-port 30001 ``` -`--tito-model` selects the tokenizer family (`qwen3`, or `default` for the -tokenizer's own template). `--backend-kind` selects the backend token dialect +`--tito-model` selects the tokenizer family (`qwen3`, `qwen3_5`, or `default` +for the tokenizer's own template); `--tito-chat-template-kwargs` pins extra +template variables as a JSON object. `--backend-kind` selects the backend token dialect (`sglang` default, or `vllm`). `--backend-url` may be omitted to auto-discover a local backend (see `agentix.tito.discovery`). Run `agentix-tito serve -h` for the full list. +### Context window clamp (`--tito-context-window TOKENS`, vllm) + +vLLM rejects any request with `prompt + max_tokens > max_model_len` — already +in `/v1/chat/completions/render`, and it counts one token more than the +session's exact prompt ids for Qwen3-family templates. For an agent whose +profile asks for a large per-turn budget (Pi with `maxTokens: 131072`) that +turns into a hard wall on trajectory length at `max_model_len - max_tokens`. +With the window configured the gateway caps the chat request's `max_tokens` +(and the rendered `sampling_params.max_tokens`) at +`context_window - len(prompt_ids) - 16` before forwarding, so the per-turn +budget shrinks as the trajectory grows instead of the turn failing with 400. +A prompt that already fills the window is left to the backend's own verdict. +Unset = forward `max_tokens` verbatim (previous behaviour). + ## Per-turn record persistence — `tito.record.v1` (normative) With `--record-dir DIR` (env `TITO_RECORD_DIR`), every committed turn appends diff --git a/plugins/tito/agentix/tito/cli.py b/plugins/tito/agentix/tito/cli.py index 340c44a..eeb24a2 100644 --- a/plugins/tito/agentix/tito/cli.py +++ b/plugins/tito/agentix/tito/cli.py @@ -3,6 +3,7 @@ from __future__ import annotations import argparse +import json import os import sys @@ -59,6 +60,16 @@ def _add_serve_arguments(parser: argparse.ArgumentParser) -> None: default=TITOTokenizerType.DEFAULT.value, help="TITO tokenizer family (qwen3, or default for the tokenizer's own template).", ) + parser.add_argument( + "--tito-chat-template-kwargs", + default=None, + metavar="JSON", + help=( + "JSON object of extra chat-template variables pinned for every render, " + 'e.g. \'{"reasoning_effort": "xhigh"}\' for Qwen3.8 (tito-model qwen3_5). ' + "The prompt tokens the model sees are rendered with exactly these values." + ), + ) parser.add_argument( "--tito-allowed-append-roles", nargs="+", @@ -93,6 +104,14 @@ def _add_serve_arguments(parser: argparse.ArgumentParser) -> None: default=None, help="LRU-evict sessions beyond this count (same flush-first, never-in-flight rules). Unset = unbounded.", ) + parser.add_argument( + "--tito-context-window", + type=int, + default=None, + metavar="TOKENS", + help="Backend context window (vLLM max_model_len). vllm turns clamp max_tokens to what fits after the " + "session's prompt ids instead of letting the backend reject the request. Unset = forward verbatim.", + ) parser.add_argument( "--backend-probe-candidate", action="append", @@ -108,6 +127,18 @@ def _add_serve_arguments(parser: argparse.ArgumentParser) -> None: ) +def _parse_template_kwargs(raw: str | None) -> dict[str, object]: + if raw is None or raw.strip() == "": + return {} + try: + parsed = json.loads(raw) + except ValueError as e: + raise SystemExit(f"--tito-chat-template-kwargs is not valid JSON: {e}") from e + if not isinstance(parsed, dict): + raise SystemExit("--tito-chat-template-kwargs must be a JSON object") + return parsed + + def _serve(args: argparse.Namespace) -> int: urls: list[str] = args.backend_url or [] config = TITOGatewayConfig.from_cli_values( @@ -120,6 +151,7 @@ def _serve(args: argparse.Namespace) -> int: chat_template_path=args.chat_template_path, tito_model=args.tito_model, tito_allowed_append_roles=args.tito_allowed_append_roles, + tito_chat_template_kwargs=_parse_template_kwargs(args.tito_chat_template_kwargs), session_server_ip=args.session_server_ip, session_server_port=args.session_server_port, router_timeout=args.router_timeout, @@ -128,6 +160,7 @@ def _serve(args: argparse.Namespace) -> int: record_dir=args.record_dir, session_ttl_seconds=args.session_ttl_seconds, max_sessions=args.max_sessions, + tito_context_window=args.tito_context_window, ) TITOGateway(config).run() return 0 diff --git a/plugins/tito/agentix/tito/config.py b/plugins/tito/agentix/tito/config.py index 2a72e74..5827a92 100644 --- a/plugins/tito/agentix/tito/config.py +++ b/plugins/tito/agentix/tito/config.py @@ -3,6 +3,7 @@ from __future__ import annotations from dataclasses import dataclass, field +from typing import Any from .discovery import DEFAULT_BACKEND_PROBE_CANDIDATES @@ -29,6 +30,10 @@ class TITOGatewayConfig: chat_template_path: str | None = None tito_model: str = "default" tito_allowed_append_roles: tuple[str, ...] = ("tool",) + # Extra chat-template variables pinned for every render (e.g. Qwen3.8 + # `reasoning_effort`). Stored as a sorted tuple of (key, value) pairs so the + # frozen config stays hashable; see `chat_template_kwargs`. + tito_chat_template_kwargs: tuple[tuple[str, Any], ...] = () session_server_ip: str = "127.0.0.1" session_server_port: int = 30000 router_timeout: float = 600.0 @@ -44,10 +49,18 @@ class TITOGatewayConfig: # touches a session with in-flight requests. session_ttl_seconds: float | None = None max_sessions: int | None = None + # Backend context window in tokens (vLLM `max_model_len`). When set, the + # vllm turn clamps the chat request's `max_tokens` to what fits after the + # session's exact prompt ids, so an agent profile with a large per-turn + # budget never trips the backend's "prompt + max_tokens > max_model_len" + # rejection as the trajectory grows. Unset = forward max_tokens verbatim. + tito_context_window: int | None = None def __post_init__(self) -> None: if not self.hf_checkpoint: raise ValueError("hf_checkpoint is required for TITO token tracking") + if self.tito_context_window is not None and self.tito_context_window <= 0: + raise ValueError("tito_context_window must be a positive token count") normalized_roles = tuple(dict.fromkeys(role.lower() for role in self.tito_allowed_append_roles)) invalid = sorted(set(normalized_roles) - _VALID_APPEND_ROLES) @@ -55,6 +68,14 @@ def __post_init__(self) -> None: raise ValueError(f"unsupported tito append roles: {invalid}") object.__setattr__(self, "tito_allowed_append_roles", normalized_roles or ("tool",)) + normalized_kwargs = tuple(sorted(dict(self.tito_chat_template_kwargs).items())) + for key, _ in normalized_kwargs: + if not isinstance(key, str) or not key: + raise ValueError("tito_chat_template_kwargs keys must be non-empty strings") + if "chat_template" in dict(normalized_kwargs): + raise ValueError("tito_chat_template_kwargs must not carry 'chat_template'; use chat_template_path") + object.__setattr__(self, "tito_chat_template_kwargs", normalized_kwargs) + if self.routing_policy not in ("sticky", "round_robin"): raise ValueError( f"routing_policy must be 'sticky' or 'round_robin'; got {self.routing_policy!r}" @@ -68,6 +89,10 @@ def __post_init__(self) -> None: if self.max_sessions is not None and self.max_sessions < 1: raise ValueError(f"max_sessions must be >= 1; got {self.max_sessions!r}") + @property + def chat_template_kwargs(self) -> dict[str, Any]: + return dict(self.tito_chat_template_kwargs) + @classmethod def from_cli_values( cls, @@ -78,6 +103,7 @@ def from_cli_values( tito_model: str, tito_allowed_append_roles: list[str], session_server_ip: str, + tito_chat_template_kwargs: dict[str, Any] | None = None, session_server_port: int, router_timeout: float, backend_urls: list[str] | None = None, @@ -89,6 +115,7 @@ def from_cli_values( record_dir: str | None = None, session_ttl_seconds: float | None = None, max_sessions: int | None = None, + tito_context_window: int | None = None, ) -> TITOGatewayConfig: return cls( hf_checkpoint=hf_checkpoint, @@ -100,6 +127,7 @@ def from_cli_values( chat_template_path=chat_template_path, tito_model=tito_model, tito_allowed_append_roles=tuple(tito_allowed_append_roles), + tito_chat_template_kwargs=tuple((tito_chat_template_kwargs or {}).items()), session_server_ip=session_server_ip, session_server_port=session_server_port, router_timeout=router_timeout, @@ -108,6 +136,7 @@ def from_cli_values( record_dir=record_dir, session_ttl_seconds=session_ttl_seconds, max_sessions=max_sessions, + tito_context_window=tito_context_window, ) def as_session_args(self): @@ -120,6 +149,7 @@ def as_session_args(self): chat_template_path=self.chat_template_path, tito_model=self.tito_model, tito_allowed_append_roles=list(self.tito_allowed_append_roles), + tito_chat_template_kwargs=self.chat_template_kwargs, trust_remote_code=self.trust_remote_code, session_server_ip=self.session_server_ip, session_server_port=self.session_server_port, @@ -127,4 +157,5 @@ def as_session_args(self): record_dir=self.record_dir, session_ttl_seconds=self.session_ttl_seconds, max_sessions=self.max_sessions, + tito_context_window=self.tito_context_window, ) diff --git a/plugins/tito/agentix/tito/engine/messages.py b/plugins/tito/agentix/tito/engine/messages.py index 8e40f46..0db62f5 100644 --- a/plugins/tito/agentix/tito/engine/messages.py +++ b/plugins/tito/agentix/tito/engine/messages.py @@ -8,6 +8,7 @@ from __future__ import annotations +import json from typing import Any # Keys a chat template actually reads. Extra client-injected keys @@ -37,8 +38,50 @@ def _reasoning_value(message: dict[str, Any]) -> Any: return None +def _canonical_tool_calls(value: Any) -> Any: + """Compare tool calls by meaning, not by serialization. + + The backend's parser emits ``function.arguments`` as one JSON string (vLLM: + ``json.dumps`` with ``": "`` spacing); an OpenAI-style client echoes it back + re-serialized (Pi/JS ``JSON.stringify``: no spaces). Both render to the same + tokens once the template parses the arguments (``normalize_tool_arguments`` + does exactly that before rendering), so they must compare equal here or an + honest second turn is mistaken for a history rewrite. Unparseable strings + are compared verbatim. + """ + normalized = normalize_value(value) + if not isinstance(normalized, list): + return normalized + result = [] + for call in normalized: + if not isinstance(call, dict): + result.append(call) + continue + function = call.get("function") + if not isinstance(function, dict): + function = {} + arguments = function.get("arguments", call.get("arguments")) + if isinstance(arguments, str): + try: + arguments = json.loads(arguments) + except ValueError: + pass + result.append( + { + "id": call.get("id"), + "name": function.get("name", call.get("name")), + "arguments": arguments, + } + ) + return result + + def message_matches(stored: dict[str, Any], new: dict[str, Any]) -> bool: for key in TEMPLATE_RELEVANT_KEYS: + if key == "tool_calls": + if _canonical_tool_calls(stored.get(key)) != _canonical_tool_calls(new.get(key)): + return False + continue if normalize_value(stored.get(key)) != normalize_value(new.get(key)): return False # Reasoning is model-generated and routinely dropped on echo (openai-python @@ -81,7 +124,7 @@ def assert_messages_append_only_with_allowed_role( f"Diffs: {diffs}" ) - for j, msg in enumerate(new_messages[len(stored_messages):]): + for j, msg in enumerate(new_messages[len(stored_messages) :]): if msg.get("role") not in allowed_append_roles: raise ValueError( f"appended message at index {len(stored_messages) + j} " diff --git a/plugins/tito/agentix/tito/engine/pretokenize.py b/plugins/tito/agentix/tito/engine/pretokenize.py index 8717e4d..05dba01 100644 --- a/plugins/tito/agentix/tito/engine/pretokenize.py +++ b/plugins/tito/agentix/tito/engine/pretokenize.py @@ -20,6 +20,7 @@ TEMPLATE_DIR = Path(__file__).parent / "templates" _VALID_ROLES = frozenset({"tool", "user", "system"}) _DUMMY_SYSTEM: dict[str, Any] = {"role": "system", "content": "dummy system"} +_DUMMY_USER: dict[str, Any] = {"role": "user", "content": "dummy user"} def _build_dummy_assistant(tool_responses: list[dict[str, Any]]) -> dict[str, Any]: @@ -233,7 +234,62 @@ def fix_prefix(self, pretokenized_token_ids: list[int]) -> list[int]: return prefix +class Qwen3_5TITOTokenizer(Qwen3TITOTokenizer): + """Qwen3.5 / Qwen3.8 (`Qwen3_5ForConditionalGeneration`): the tokenizer's OWN + chat template, not a bundled fixed one. + + Two properties of that template family shape this subclass (verified on the + real Qwen/Qwen3.8-27B tokenizer, 2026-09-02): + + - it raises ``No user query found in messages`` when the conversation has + no real user turn, so the synthetic contexts the incremental algorithm + renders (`dummy system` + `dummy assistant`) must also carry a dummy + user message — the suffix diff is unaffected because the dummy turns are + in both renders; + - an appended ``system`` message raises ``System message must be at the + beginning``, so only ``tool`` (and, for Qwen3.8 whose template defaults + to ``preserve_thinking``, ``user``) appends are supported. + + Reasoning is rendered from ``reasoning_content`` only (the `` + content-splitting fallback of Qwen3 is gone), tool calls use the + ```` XML dialect (vLLM parser ``qwen3_coder``), + and the model stops at ``<|im_end|>`` without the template's trailing + newline exactly like Qwen3, so ``fix_prefix`` is inherited. + + ``chat_template_kwargs`` is where the caller pins template variables that + change the rendered prompt — for Qwen3.8 ``reasoning_effort`` + (``xhigh``/``medium``/``low``; the template default is ``xhigh``) and + ``preserve_thinking``. They are part of the tokenizer identity the record + carries via ``chat_template_sha256`` only indirectly, so a gateway serving + one experiment must run with one fixed set. + """ + + reasoning_parser = "qwen3" + tool_call_parser = "qwen3_coder" + + def _synthetic_base(self) -> list[dict[str, Any]]: + return [_DUMMY_SYSTEM, _DUMMY_USER] + + def _tokenize_tool_segment( + self, appended_messages: list[dict[str, Any]], tools: list[dict[str, Any]] | None = None + ) -> list[int]: + return self._tokenize_rendered_suffix( + [*self._synthetic_base(), _build_dummy_assistant(appended_messages)], appended_messages, tools=tools + ) + + def _tokenize_user_and_system_segment( + self, appended_message: dict[str, Any], tools: list[dict[str, Any]] | None = None + ) -> list[int]: + if appended_message.get("role") == "system": + raise ValueError( + "the Qwen3.5/3.8 chat template only accepts a system message at position 0; " + "appending one mid-conversation cannot be tokenized incrementally" + ) + return self._tokenize_rendered_suffix(self._synthetic_base(), [appended_message], tools=tools) + + _QWEN3_FIXED = "qwen3_fixed.jinja" +SUPPORTED_TOKENIZER_TYPES = ("default", "qwen3", "qwen3_5") def get_tito_tokenizer( @@ -241,22 +297,36 @@ def get_tito_tokenizer( tokenizer_type: str = "qwen3", *, allowed_append_roles: tuple[str, ...] = ("tool",), + chat_template_kwargs: dict[str, Any] | None = None, ) -> TITOTokenizer: """Build a TITO tokenizer. `default` uses the tokenizer's own chat template (model-agnostic); `qwen3` loads the bundled fixed template (and disables thinking - clearing when `user` appends are allowed, so earlier turns keep their reasoning).""" + clearing when `user` appends are allowed, so earlier turns keep their reasoning); + `qwen3_5` uses the tokenizer's own Qwen3.5/3.8 template with the synthetic-context + fix that family needs. `chat_template_kwargs` are extra template variables pinned + for every render (e.g. Qwen3.8 ``reasoning_effort``); for `qwen3` they merge on + top of the fixed-template override.""" if tokenizer is None: raise ValueError("tokenizer must not be None") roles = frozenset(allowed_append_roles) invalid = roles - _VALID_ROLES if invalid: raise ValueError(f"unknown roles in allowed_append_roles: {sorted(invalid)}; valid: {sorted(_VALID_ROLES)}") + extra = dict(chat_template_kwargs or {}) + if "chat_template" in extra: + raise ValueError("chat_template_kwargs must not carry 'chat_template'; use --chat-template-path") if tokenizer_type == "default": - return TITOTokenizer(tokenizer, allowed_append_roles=list(allowed_append_roles)) + return TITOTokenizer(tokenizer, chat_template_kwargs=extra, allowed_append_roles=list(allowed_append_roles)) if tokenizer_type == "qwen3": kw: dict[str, Any] = {"chat_template": (TEMPLATE_DIR / _QWEN3_FIXED).read_text()} if "user" in roles: kw["clear_thinking"] = False + kw.update(extra) return Qwen3TITOTokenizer(tokenizer, chat_template_kwargs=kw, allowed_append_roles=list(allowed_append_roles)) - raise ValueError(f"unsupported tokenizer_type {tokenizer_type!r}; supported: 'qwen3', 'default'") + if tokenizer_type == "qwen3_5": + return Qwen3_5TITOTokenizer( + tokenizer, chat_template_kwargs=extra, allowed_append_roles=list(allowed_append_roles) + ) + supported = ", ".join(repr(t) for t in SUPPORTED_TOKENIZER_TYPES) + raise ValueError(f"unsupported tokenizer_type {tokenizer_type!r}; supported: {supported}") diff --git a/plugins/tito/agentix/tito/engine/render.py b/plugins/tito/agentix/tito/engine/render.py index ee69863..cd5f5fd 100644 --- a/plugins/tito/agentix/tito/engine/render.py +++ b/plugins/tito/agentix/tito/engine/render.py @@ -37,6 +37,18 @@ def normalize_tool_arguments(messages: list[dict], format: Literal["dict", "json if msg.get("role") == "assistant": if msg.get("content") is None: msg["content"] = "" + # Clients echo reasoning under the key their provider face used: + # vLLM emits `reasoning`, Pi echoes `reasoning` back, while Qwen + # templates only read `reasoning_content`. Without this alias a + # from-scratch render drops every earlier turn's thinking + # (empty ), diverging from the prefix the model actually + # generated on. + if not msg.get("reasoning_content"): + for key in ("reasoning", "reasoning_text"): + value = msg.get(key) + if isinstance(value, str) and value: + msg["reasoning_content"] = value + break if isinstance(msg.get("tool_calls"), list): for item in msg["tool_calls"]: func = item.get("function") diff --git a/plugins/tito/agentix/tito/engine/session_app.py b/plugins/tito/agentix/tito/engine/session_app.py index 4a7f287..a8ffa2a 100644 --- a/plugins/tito/agentix/tito/engine/session_app.py +++ b/plugins/tito/agentix/tito/engine/session_app.py @@ -32,6 +32,7 @@ from .pretokenize import get_tito_tokenizer from .processing import load_tokenizer from .record import build_turn_record, sampling_from_request +from .sse import build_sse_response from .trajectory import GetSessionResponse, LinearTrajectory, SessionRecord, SessionRegistry from .upstream import Backend, get_upstream @@ -56,6 +57,7 @@ def build_registry(args: Any) -> SessionRegistry | None: tokenizer, tokenizer_type=getattr(args, "tito_model", "default"), allowed_append_roles=tuple(roles), + chat_template_kwargs=dict(getattr(args, "tito_chat_template_kwargs", None) or {}), ) return SessionRegistry(args, tokenizer, tito_tokenizer=tito_tokenizer) @@ -70,7 +72,12 @@ def setup_session_routes(app: FastAPI, backend: Backend, args: Any) -> None: # capture must be complete and closed however the process exits cleanly. app.router.on_shutdown.append(registry.close) - adapter = get_upstream(getattr(args, "backend_kind", "sglang")) + adapter = get_upstream( + getattr(args, "backend_kind", "sglang"), + context_window=getattr(args, "tito_context_window", None), + eos_token_id=getattr(registry.tokenizer, "eos_token_id", None), + eos_token=getattr(registry.tokenizer, "eos_token", None), + ) backend_kind = str(getattr(args, "backend_kind", "sglang") or "sglang") instance_id = getattr(args, "session_server_instance_id", None) @@ -157,6 +164,11 @@ async def _chat_turn(request: Request, session_id: str, session: LinearTrajector # Adapter preconditions raise HERE, before the lock: phase 1 can # commit a rollback, and a rejected request must leave no side effects. adapter.validate_request(request_body) + # The adapter forces stream=false upstream (the token harvest needs the + # whole JSON body). Remember what the CLIENT asked for: an agent that + # only speaks SSE (Pi, the OpenAI SDK with stream=True) must get its + # completed turn back as a well-formed event stream, not a JSON body. + client_wants_stream = request_body.get("stream") is True # Phase 1: prepare the pretokenized prompt ids (lock held briefly). async with session.lock: @@ -176,6 +188,8 @@ async def _chat_turn(request: Request, session_id: str, session: LinearTrajector # break token accumulation) and harvests the exact completion ids. turn = await adapter.chat_turn(backend, request, request_body, prompt_token_ids) if turn.harvest is None: + # Upstream non-200: pass the error body through as-is (an SSE + # client treats a non-2xx as an error regardless of framing). return backend.build_proxy_response(turn.proxy_result) harvest = turn.harvest @@ -241,6 +255,8 @@ async def _chat_turn(request: Request, session_id: str, session: LinearTrajector ) except Exception: logger.exception("tito record: failed to capture turn for session %s", session_id) + if client_wants_stream: + return build_sse_response(turn.proxy_result["response_body"]) return backend.build_proxy_response(turn.proxy_result) @app.api_route("/sessions/{session_id}/{path:path}", methods=["GET", "POST", "PUT", "DELETE", "PATCH"]) diff --git a/plugins/tito/agentix/tito/engine/sse.py b/plugins/tito/agentix/tito/engine/sse.py new file mode 100644 index 0000000..ec4186f --- /dev/null +++ b/plugins/tito/agentix/tito/engine/sse.py @@ -0,0 +1,90 @@ +"""Re-frame a completed chat completion as an OpenAI-style SSE stream. + +The TITO flow forces ``stream=false`` upstream because the token harvest needs +the whole JSON body (``meta_info`` / logprobs). Agents that only speak the +streaming dialect — Pi's ``openai-completions`` provider always sends +``stream: true`` and parses Server-Sent Events, raising "Stream ended without +finish_reason" on anything else — would otherwise break at their first turn. + +This module turns the finished completion into the minimal event stream such +a client accepts: one chunk carrying the whole assistant delta (role, content, +reasoning under whatever key the backend used, tool calls with their ``index``), +one terminal chunk with ``finish_reason`` and ``usage``, then ``[DONE]``. The +content is byte-identical to what the JSON body carried; only the framing +changes, so the recorded turn and the served turn stay the same tokens. +""" + +from __future__ import annotations + +import json +from typing import Any + +from starlette.responses import Response + +REASONING_KEYS = ("reasoning_content", "reasoning", "reasoning_text") + + +def completion_to_chunks(completion: dict[str, Any]) -> list[dict[str, Any]]: + """The chunk objects (already ordered) for one non-streamed completion.""" + choices = completion.get("choices") + if not isinstance(choices, list) or not choices or not isinstance(choices[0], dict): + raise ValueError("completion has no choices to stream") + choice = choices[0] + message = choice.get("message") + if not isinstance(message, dict): + raise ValueError("completion choice has no message to stream") + + base: dict[str, Any] = { + "id": completion.get("id"), + "object": "chat.completion.chunk", + "created": completion.get("created"), + "model": completion.get("model"), + } + delta: dict[str, Any] = {"role": message.get("role") or "assistant"} + content = message.get("content") + if isinstance(content, str) and content: + delta["content"] = content + for key in REASONING_KEYS: + value = message.get(key) + if isinstance(value, str) and value: + delta[key] = value + tool_calls = message.get("tool_calls") + if isinstance(tool_calls, list) and tool_calls: + delta["tool_calls"] = [ + {"index": index, **call} if isinstance(call, dict) and "index" not in call else call + for index, call in enumerate(tool_calls) + ] + first = {**base, "choices": [{"index": choice.get("index", 0), "delta": delta, "finish_reason": None}]} + final: dict[str, Any] = { + **base, + "choices": [ + { + "index": choice.get("index", 0), + "delta": {}, + "finish_reason": choice.get("finish_reason") or "stop", + } + ], + } + if completion.get("usage") is not None: + final["usage"] = completion["usage"] + return [first, final] + + +def render_sse(chunks: list[dict[str, Any]]) -> bytes: + body = "".join(f"data: {json.dumps(chunk, ensure_ascii=False)}\n\n" for chunk in chunks) + return (body + "data: [DONE]\n\n").encode("utf-8") + + +def build_sse_response(response_body: bytes) -> Response: + """SSE response for a 200 JSON completion body the harvest already parsed. + + The body was validated by the upstream adapter (choices/message present), + so a failure here is a programming error, not a client-facing 4xx. + """ + completion = json.loads(response_body) + return Response( + content=render_sse(completion_to_chunks(completion)), + status_code=200, + media_type="text/event-stream", + headers={"cache-control": "no-cache", "x-tito-stream": "reframed"}, + ) diff --git a/plugins/tito/agentix/tito/engine/upstream.py b/plugins/tito/agentix/tito/engine/upstream.py index dc9777c..ffd3e73 100644 --- a/plugins/tito/agentix/tito/engine/upstream.py +++ b/plugins/tito/agentix/tito/engine/upstream.py @@ -107,11 +107,17 @@ def _extract_assistant_message(choice: dict) -> dict: assistant_message = choice.get("message") if not isinstance(assistant_message, dict): raise UpstreamResponseError("assistant message missing") - if assistant_message.get("content") is None and not assistant_message.get("tool_calls"): + has_reasoning = any( + isinstance(assistant_message.get(key), str) and assistant_message.get(key) + for key in ("reasoning_content", "reasoning", "reasoning_text") + ) + if assistant_message.get("content") is None and not assistant_message.get("tool_calls") and not has_reasoning: # Tool-call-only turns routinely carry content:null (the parser - # consumed all generated text) — only a turn with NEITHER content - # NOR tool_calls is malformed. - raise UpstreamResponseError("assistant message has neither content nor tool_calls") + # consumed all generated text), and a reasoning model that hits + # max_tokens inside its block yields reasoning with no visible + # content (finish_reason=length) — both are real, recordable turns. + # Only a turn with NONE of content / tool_calls / reasoning is malformed. + raise UpstreamResponseError("assistant message has neither content, tool_calls, nor reasoning") return assistant_message @@ -128,9 +134,9 @@ async def chat_turn( self, backend: Backend, request: Request, request_body: dict, prompt_token_ids: list[int] ) -> ChatTurn: # Hardcoded so an agent override can't break token accumulation: - request_body["logprobs"] = True # -> meta_info.output_token_logprobs - request_body["return_meta_info"] = True # -> choice.meta_info - request_body["no_stop_trim"] = False # stop-token text trimmed from content + request_body["logprobs"] = True # -> meta_info.output_token_logprobs + request_body["return_meta_info"] = True # -> choice.meta_info + request_body["no_stop_trim"] = False # stop-token text trimmed from content # The TITO flow needs the complete JSON completion (logprobs + # meta_info); an SSE stream would be unparseable below. Force # non-streaming — a stream:true agent gets the full JSON body back. @@ -206,10 +212,82 @@ class VllmUpstream: release with the derender endpoints).""" kind = "vllm" + # vLLM's render re-tokenizes the message history and, for the Qwen3 family, + # counts one token more than the session's exact prompt ids (the + # `<|im_end|>\n` fixup); its `prompt + max_tokens <= max_model_len` check + # runs against THAT count. Leave a little room so an exact fit is not + # rejected by an off-by-one. + CONTEXT_MARGIN = 16 + + def __init__( + self, context_window: int | None = None, eos_token_id: int | None = None, eos_token: str | None = None + ) -> None: + self.context_window = context_window + self.eos_token_id = eos_token_id + self.eos_token = eos_token + + def normalize_terminal_eos(self, response: dict, token_ids: list[int]) -> bool: + choice = _first_choice(response) + if ( + not self.eos_token + or self.eos_token_id is None + or not token_ids + or token_ids[-1] != self.eos_token_id + or choice.get("finish_reason") != "stop" + ): + return False + message = choice.get("message") + if not isinstance(message, dict): + return False + for key in ("content", "reasoning", "reasoning_content"): + text = message.get(key) + if isinstance(text, str) and text: + if text.endswith(self.eos_token): + message[key] = text[: -len(self.eos_token)] + return True + return False + return False def validate_request(self, request_body: dict) -> None: _require_model(request_body) + def _room(self, prompt_len: int) -> int | None: + if self.context_window is None: + return None + room = self.context_window - prompt_len - self.CONTEXT_MARGIN + return room if room > 0 else None + + def clamp_chat_max_tokens(self, request_body: dict, prompt_len: int) -> int | None: + """Fit the chat request's ``max_tokens`` into the backend's context window. + + vLLM validates ``prompt + max_tokens <= max_model_len`` already in + ``/v1/chat/completions/render`` (and again in generate), which turns a + generous per-turn budget into a hard wall on trajectory length. The + gateway knows the session's exact prompt length before render, so cap + the budget at the remaining room here. Returns the effective value + (None when nothing changed). A prompt that already fills the window is + left alone: the backend's own rejection is the right verdict there. + """ + room = self._room(prompt_len) + requested = request_body.get("max_tokens") + if room is None or not isinstance(requested, int) or requested <= room: + return None + request_body["max_tokens"] = room + return room + + def clamp_max_tokens(self, generate_request: dict, prompt_len: int) -> int | None: + """Same fit on the rendered ``GenerateRequest.sampling_params`` (render + may have resolved a server-side default larger than the room).""" + room = self._room(prompt_len) + params = generate_request.get("sampling_params") + if room is None or not isinstance(params, dict): + return None + requested = params.get("max_tokens") + if not isinstance(requested, int) or requested <= room: + return None + params["max_tokens"] = room + return room + async def chat_turn( self, backend: Backend, request: Request, request_body: dict, prompt_token_ids: list[int] ) -> ChatTurn: @@ -227,10 +305,9 @@ async def chat_turn( # derender only accepts a complete (non-streamed) GenerateResponse. request_body["stream"] = False request_body.pop("stream_options", None) + self.clamp_chat_max_tokens(request_body, len(prompt_token_ids)) - render_result = await backend.do_proxy( - request, "v1/chat/completions/render", body=_encode(request_body) - ) + render_result = await backend.do_proxy(request, "v1/chat/completions/render", body=_encode(request_body)) if render_result["status_code"] != 200: return ChatTurn(proxy_result=render_result) generate_request = _parse_response_object(render_result["response_body"]) @@ -246,10 +323,9 @@ async def chat_turn( # render copies `stream` from the chat request it saw — re-force. generate_request["stream"] = False generate_request.pop("stream_options", None) + self.clamp_max_tokens(generate_request, len(prompt_token_ids)) - generate_result = await backend.do_proxy( - request, "inference/v1/generate", body=_encode(generate_request) - ) + generate_result = await backend.do_proxy(request, "inference/v1/generate", body=_encode(generate_request)) if generate_result["status_code"] != 200: return ChatTurn(proxy_result=generate_result) generate_response = _parse_response_object(generate_result["response_body"]) @@ -272,6 +348,7 @@ async def chat_turn( if derender_result["status_code"] != 200: return ChatTurn(proxy_result=derender_result) response = _parse_response_object(derender_result["response_body"]) + normalized_eos = self.normalize_terminal_eos(response, completion_token_ids) assistant_message = _extract_assistant_message(_first_choice(response)) if assistant_message.get("reasoning") is not None and assistant_message.get("reasoning_content") is None: # The engine's templates and the mismatch audit read the @@ -279,7 +356,7 @@ async def chat_turn( # trajectory copy only — the wire response keeps vLLM's # `reasoning` verbatim. assistant_message = {**assistant_message, "reasoning_content": assistant_message["reasoning"]} - if _rewrite_tool_call_finish_reasons(response): + if _rewrite_tool_call_finish_reasons(response) or normalized_eos: derender_result = {**derender_result, "response_body": _encode(response)} return ChatTurn( @@ -318,9 +395,7 @@ def _harvest_generate_tokens(generate_response: dict) -> tuple[list[int], list[f # turn without the per-token cross-check. raise UpstreamResponseError("generate response logprobs missing (needs logprobs forced on)") if len(content) != len(token_ids): - raise UpstreamResponseError( - f"len(logprobs.content)={len(content)} != len(token_ids)={len(token_ids)}" - ) + raise UpstreamResponseError(f"len(logprobs.content)={len(content)} != len(token_ids)={len(token_ids)}") try: completion_logprobs = [float(entry["logprob"]) for entry in content] except (TypeError, ValueError, KeyError) as e: @@ -328,9 +403,11 @@ def _harvest_generate_tokens(generate_response: dict) -> tuple[list[int], list[f return list(token_ids), completion_logprobs -def get_upstream(kind: str) -> UpstreamAdapter: +def get_upstream( + kind: str, *, context_window: int | None = None, eos_token_id: int | None = None, eos_token: str | None = None +) -> UpstreamAdapter: if kind == "sglang": return SglangUpstream() if kind == "vllm": - return VllmUpstream() + return VllmUpstream(context_window=context_window, eos_token_id=eos_token_id, eos_token=eos_token) raise ValueError(f"unsupported backend_kind {kind!r}; supported: {list(BACKEND_KINDS)}") diff --git a/plugins/tito/agentix/tito/tokenizer.py b/plugins/tito/agentix/tito/tokenizer.py index bf3fa11..d5854e6 100644 --- a/plugins/tito/agentix/tito/tokenizer.py +++ b/plugins/tito/agentix/tito/tokenizer.py @@ -14,6 +14,7 @@ class TITOTokenizerType(StrEnum): DEFAULT = "default" QWEN3 = "qwen3" + QWEN3_5 = "qwen3_5" def get_tito_tokenizer( @@ -21,9 +22,12 @@ def get_tito_tokenizer( tokenizer_type: TITOTokenizerType | str = TITOTokenizerType.DEFAULT, *, allowed_append_roles: tuple[str, ...] | list[str] | None = None, + chat_template_kwargs: dict[str, Any] | None = None, **_ignored: Any, ) -> Any: - """Build a TITO tokenizer for *tokenizer* (`"qwen3"` or `"default"`).""" + """Build a TITO tokenizer for *tokenizer* (`"qwen3"`, `"qwen3_5"`, or `"default"`).""" t = tokenizer_type.value if isinstance(tokenizer_type, TITOTokenizerType) else str(tokenizer_type) roles = tuple(allowed_append_roles) if allowed_append_roles else ("tool",) - return _engine_get_tito_tokenizer(tokenizer, t, allowed_append_roles=roles) + return _engine_get_tito_tokenizer( + tokenizer, t, allowed_append_roles=roles, chat_template_kwargs=chat_template_kwargs + ) diff --git a/plugins/tito/tests/package/test_cli.py b/plugins/tito/tests/package/test_cli.py index bbdb003..7257649 100644 --- a/plugins/tito/tests/package/test_cli.py +++ b/plugins/tito/tests/package/test_cli.py @@ -33,13 +33,29 @@ def test_cli_serve_parses_args(): assert args.tito_allowed_append_roles == ["tool", "user"] -def test_cli_tito_model_choices_are_qwen3_and_default(): +def test_cli_tito_model_choices_are_qwen3_qwen3_5_and_default(): args = build_parser().parse_args(["serve", "--hf-checkpoint", "X"]) assert args.tito_model == "default" + args = build_parser().parse_args(["serve", "--hf-checkpoint", "X", "--tito-model", "qwen3_5"]) + assert args.tito_model == "qwen3_5" with pytest.raises(SystemExit): build_parser().parse_args(["serve", "--hf-checkpoint", "X", "--tito-model", "glm47"]) +def test_cli_chat_template_kwargs_parse_as_json_object(): + from agentix.tito.cli import _parse_template_kwargs + + args = build_parser().parse_args( + ["serve", "--hf-checkpoint", "X", "--tito-chat-template-kwargs", '{"reasoning_effort": "xhigh"}'] + ) + assert _parse_template_kwargs(args.tito_chat_template_kwargs) == {"reasoning_effort": "xhigh"} + assert _parse_template_kwargs(None) == {} + with pytest.raises(SystemExit): + _parse_template_kwargs("[1]") + with pytest.raises(SystemExit): + _parse_template_kwargs("{not json") + + def test_cli_backend_url_is_repeatable_for_a_pool(): args = build_parser().parse_args( ["serve", "--hf-checkpoint", "X", diff --git a/plugins/tito/tests/package/test_engine.py b/plugins/tito/tests/package/test_engine.py index f9f451b..948d56f 100644 --- a/plugins/tito/tests/package/test_engine.py +++ b/plugins/tito/tests/package/test_engine.py @@ -12,7 +12,7 @@ from agentix.tito.engine.compare import MismatchType, TokenSeqComparator from agentix.tito.engine.errors import TokenizationError from agentix.tito.engine.messages import assert_messages_append_only_with_allowed_role, message_matches -from agentix.tito.engine.pretokenize import Qwen3TITOTokenizer, get_tito_tokenizer +from agentix.tito.engine.pretokenize import Qwen3_5TITOTokenizer, Qwen3TITOTokenizer, get_tito_tokenizer from agentix.tito.engine.trajectory import LinearTrajectory, SessionRegistry from tokenizers import Tokenizer, models, pre_tokenizers from transformers import PreTrainedTokenizerFast @@ -74,6 +74,52 @@ def test_message_matches_collapses_falsy_sentinels(): assert not message_matches({"role": "u", "content": "x"}, {"role": "t", "content": "x"}) +def test_render_aliases_reasoning_key_to_reasoning_content(): + """Pi echoes earlier assistant turns with `reasoning` (vLLM's key); Qwen + templates read `reasoning_content`. Rendering must treat them the same.""" + from agentix.tito.engine.render import normalize_tool_arguments + + via_reasoning = normalize_tool_arguments( + [{"role": "assistant", "content": "ok", "reasoning": "because"}], "dict" + )[0] + assert via_reasoning["reasoning_content"] == "because" + kept = normalize_tool_arguments( + [{"role": "assistant", "content": "ok", "reasoning": "x", "reasoning_content": "keep me"}], "dict" + )[0] + assert kept["reasoning_content"] == "keep me" + untouched = normalize_tool_arguments([{"role": "user", "content": "hi", "reasoning": "n/a"}], "dict")[0] + assert "reasoning_content" not in untouched + + +def test_message_matches_compares_tool_call_arguments_by_meaning(): + """vLLM stores arguments as json.dumps (": " spacing); Pi echoes them back + via JSON.stringify (no spaces). Same call, different bytes, must match — + otherwise the second turn of every tool-using rollout is a false rollback.""" + stored = {"role": "assistant", "content": "", "tool_calls": [{ + "id": "c1", "type": "function", + "function": {"name": "bash", "arguments": "{\"command\": \"ls -la && cat reference.py\"}"}, + }]} + echoed = {"role": "assistant", "content": None, "tool_calls": [{ + "id": "c1", "type": "function", + "function": {"name": "bash", "arguments": "{\"command\":\"ls -la && cat reference.py\"}"}, + }]} + assert message_matches(stored, echoed) + # A genuinely different call is still a mismatch. + other = {"role": "assistant", "content": None, "tool_calls": [{ + "id": "c1", "type": "function", + "function": {"name": "bash", "arguments": "{\"command\":\"rm -rf /\"}"}, + }]} + assert not message_matches(stored, other) + # Different id or name also mismatch; unparseable strings compare verbatim. + assert not message_matches(stored, {**echoed, "tool_calls": [{**echoed["tool_calls"][0], "id": "c2"}]}) + def raw(arguments): + call = {"id": "c1", "function": {"name": "f", "arguments": arguments}} + return {"role": "assistant", "content": "", "tool_calls": [call]} + + raw_a, raw_b = raw("not json"), raw("not json") + assert not message_matches(raw_a, raw_b) + + def test_message_matches_reasoning_tolerates_absence_and_key_dialect(): """Reasoning is model-generated and routinely dropped on echo (openai-python keeps only the standard fields), and lives under @@ -305,3 +351,82 @@ def test_version_advances_on_rollback_and_update(tok): ) assert tr.num_assistant == n0 assert tr.version != v0 + + +# A Qwen3.5-shaped template on the tiny vocab: raises without a real user +# turn, wraps tool results in a user turn, and writes `<|im_end|>` + newline +# after every message — the properties the qwen3_5 family exists for. +_QWEN35_LIKE_TEMPLATE = ( + "{%- set ns = namespace(found=false) -%}" + "{%- for m in messages -%}{%- if m['role'] == 'user' -%}{%- set ns.found = true -%}{%- endif -%}{%- endfor -%}" + "{%- if not ns.found -%}{{ raise_exception('No user query found in messages.') }}{%- endif -%}" + "{%- for m in messages -%}" + "{%- if m['role'] == 'system' and not loop.first -%}" + "{{ raise_exception('System message must be at the beginning.') }}{%- endif -%}" + "{%- if m['role'] == 'tool' -%}<|im_start|>user {{ m['content'] or '' }}<|im_end|>{{ '\\n' }}" + "{%- else -%}<|im_start|>{{ m['role'] }} {{ m['content'] or '' }}<|im_end|>{{ '\\n' }}{%- endif -%}" + "{%- endfor -%}" + "{%- if add_generation_prompt -%}<|im_start|>assistant {%- endif -%}" +) + + +@pytest.fixture(scope="module") +def qwen35_like_tok(): + specials = ["", "", "", "<|im_start|>", "<|im_end|>"] + words = ["system", "user", "assistant", "tool", "dummy", "You", "are", "ok", + "done", "compute", "17", "23", "391", "X", "Y", "Hello", "\n"] + vocab = {t: i for i, t in enumerate(specials + words)} + tk = Tokenizer(models.WordLevel(vocab=vocab, unk_token="")) + tk.pre_tokenizer = pre_tokenizers.WhitespaceSplit() + t = PreTrainedTokenizerFast( + tokenizer_object=tk, unk_token="", bos_token="", eos_token="", + additional_special_tokens=["<|im_start|>", "<|im_end|>"], + ) + # The Qwen families need "\n" to be one token (the `<|im_end|>` newline fixup). + t.add_tokens(["\n"]) + t.chat_template = _QWEN35_LIKE_TEMPLATE + return t + + +def test_qwen3_5_family_renders_synthetic_contexts_with_a_user_turn(qwen35_like_tok): + tok = qwen35_like_tok + # The `default` engine's synthetic base (dummy system only) is rejected by + # this template family ... + default = get_tito_tokenizer(tok, "default", allowed_append_roles=("tool",)) + with pytest.raises(ValueError): + default.tokenize_additional_non_assistant( + [{"role": "user", "content": "Hello"}, {"role": "assistant", "content": "ok"}], + [{"role": "user", "content": "Hello"}, {"role": "assistant", "content": "ok"}, + {"role": "tool", "content": "17", "tool_call_id": "c1"}], + ) + # ... while qwen3_5 carries a dummy user turn and the incremental result + # equals the from-scratch render (the engine invariant). + tt = get_tito_tokenizer(tok, "qwen3_5", allowed_append_roles=("tool",)) + assert isinstance(tt, Qwen3_5TITOTokenizer) + assert tt.tool_call_parser == "qwen3_coder" + old = [{"role": "system", "content": "You are"}, {"role": "user", "content": "Hello"}, + {"role": "assistant", "content": "ok"}] + new = old + [{"role": "tool", "content": "17", "tool_call_id": "c1"}] + prefix = tt.render_messages(old, add_generation_prompt=False, tokenize=True) + merged = tt.merge_tokens(old, new, prefix) + assert merged == tt.render_messages(new, add_generation_prompt=True, tokenize=True) + + +def test_qwen3_5_family_rejects_mid_conversation_system_append(qwen35_like_tok): + tt = get_tito_tokenizer(qwen35_like_tok, "qwen3_5", allowed_append_roles=("tool", "system")) + old = [{"role": "user", "content": "Hello"}, {"role": "assistant", "content": "ok"}] + with pytest.raises(ValueError, match="position 0"): + tt.tokenize_additional_non_assistant(old, old + [{"role": "system", "content": "X"}]) + + +def test_chat_template_kwargs_are_pinned_into_every_render(qwen35_like_tok): + tok = qwen35_like_tok + tok.chat_template = "{{ effort }}:" + _QWEN35_LIKE_TEMPLATE + try: + tt = get_tito_tokenizer(tok, "qwen3_5", chat_template_kwargs={"effort": "X"}) + text = tt.render_messages([{"role": "user", "content": "Hello"}], add_generation_prompt=True) + assert text.startswith("X:") + with pytest.raises(ValueError, match="chat_template"): + get_tito_tokenizer(tok, "qwen3_5", chat_template_kwargs={"chat_template": "y"}) + finally: + tok.chat_template = _QWEN35_LIKE_TEMPLATE diff --git a/plugins/tito/tests/test_gateway_http.py b/plugins/tito/tests/test_gateway_http.py index e4d8948..615c99f 100644 --- a/plugins/tito/tests/test_gateway_http.py +++ b/plugins/tito/tests/test_gateway_http.py @@ -155,22 +155,82 @@ async def test_full_session_flow_over_http(gateway): assert r.status_code == 404 +def _sse_chunks(body: str) -> list: + events = [line[len("data: "):] for line in body.split("\n") if line.startswith("data: ")] + assert events[-1] == "[DONE]" + return [json.loads(e) for e in events[:-1]] + + @pytest.mark.asyncio -async def test_stream_request_is_forced_non_streaming(gateway): +async def test_stream_request_is_forced_non_streaming_upstream_but_served_as_sse(gateway): """The TITO flow needs the full JSON completion (logprobs + meta_info), so - the gateway must force stream=false upstream and answer 200 with the JSON - body — not 500 on an unparseable SSE stream.""" + the gateway forces stream=false upstream; the CLIENT asked for a stream, + so it gets the completed turn re-framed as SSE (Pi and the OpenAI SDK + reject a JSON body when they asked for stream=true).""" client, replica, _ = gateway + replica.message = { + "role": "assistant", + "content": "ok done", + "reasoning_content": "think", + "tool_calls": [{ + "id": "call_1", "type": "function", + "function": {"name": "compute", "arguments": "{\"x\": 1}"}, + }], + } sid = (await client.post("/sessions")).json()["session_id"] r = await client.post( f"/sessions/{sid}/v1/chat/completions", json={**_CHAT, "stream": True} ) assert r.status_code == 200 - assert r.json()["choices"][0]["message"]["content"] == "ok done" + assert r.headers["content-type"].startswith("text/event-stream") _, seen = replica.calls[0] assert seen["stream"] is False + first, last = _sse_chunks(r.text) + assert first["object"] == "chat.completion.chunk" + assert first["id"] == "c1" + delta = first["choices"][0]["delta"] + assert delta["role"] == "assistant" + assert delta["content"] == "ok done" + assert delta["reasoning_content"] == "think" + assert delta["tool_calls"] == [{ + "index": 0, "id": "call_1", "type": "function", + "function": {"name": "compute", "arguments": "{\"x\": 1}"}, + }] + assert first["choices"][0]["finish_reason"] is None + assert last["choices"][0]["finish_reason"] == "stop" + assert last["usage"] == replica.usage + + # The turn was recorded exactly as if the client had not streamed. + got = (await client.get(f"/sessions/{sid}")).json() + assert len(got["records"]) == 1 + assert got["metadata"]["accumulated_token_ids"][-2:] == [7, 8] + + +@pytest.mark.asyncio +async def test_non_streaming_request_still_gets_the_json_body(gateway): + client, replica, _ = gateway + sid = (await client.post("/sessions")).json()["session_id"] + r = await client.post(f"/sessions/{sid}/v1/chat/completions", json=_CHAT) + assert r.status_code == 200 + assert r.headers["content-type"].startswith("application/json") + assert r.json()["choices"][0]["message"]["content"] == "ok done" + + +@pytest.mark.asyncio +async def test_stream_request_upstream_error_passes_through_as_json(gateway): + """A non-200 upstream is not re-framed: SSE clients treat any non-2xx as + an error, and the body must stay the backend's verbatim message.""" + client, replica, _ = gateway + replica.message = {"role": "assistant", "content": None} + sid = (await client.post("/sessions")).json()["session_id"] + r = await client.post( + f"/sessions/{sid}/v1/chat/completions", json={**_CHAT, "stream": True} + ) + assert r.status_code == 502 + assert not r.headers["content-type"].startswith("text/event-stream") + @pytest.mark.asyncio async def test_tool_call_completion_with_null_content_is_accepted(gateway): @@ -194,6 +254,21 @@ async def test_tool_call_completion_with_null_content_is_accepted(gateway): assert len(got["records"]) == 1 # the turn was recorded +@pytest.mark.asyncio +async def test_reasoning_only_truncated_turn_is_accepted_and_recorded(gateway): + """A thinking model cut off by max_tokens inside returns + reasoning_content with content:null and no tool_calls (vLLM derender, + finish_reason=length). That is a real generation the trajectory must keep.""" + client, replica, _ = gateway + replica.message = {"role": "assistant", "content": None, "reasoning_content": "still thinking"} + sid = (await client.post("/sessions")).json()["session_id"] + r = await client.post(f"/sessions/{sid}/v1/chat/completions", json=_CHAT) + assert r.status_code == 200 + assert r.json()["choices"][0]["message"]["reasoning_content"] == "still thinking" + got = (await client.get(f"/sessions/{sid}")).json() + assert len(got["records"]) == 1 + + @pytest.mark.asyncio async def test_content_none_without_tool_calls_is_still_502(gateway): client, replica, _ = gateway diff --git a/plugins/tito/tests/test_gateway_vllm_http.py b/plugins/tito/tests/test_gateway_vllm_http.py index 3ac1109..ff3b49a 100644 --- a/plugins/tito/tests/test_gateway_vllm_http.py +++ b/plugins/tito/tests/test_gateway_vllm_http.py @@ -35,7 +35,10 @@ def tok(): tk = Tokenizer(models.WordLevel(vocab=vocab, unk_token="")) tk.pre_tokenizer = pre_tokenizers.Whitespace() t = PreTrainedTokenizerFast( - tokenizer_object=tk, unk_token="", bos_token="", eos_token="", + tokenizer_object=tk, + unk_token="", + bos_token="", + eos_token="", additional_special_tokens=["<|im_start|>", "<|im_end|>"], ) t.chat_template = ( @@ -54,6 +57,7 @@ def _args(): tito_model="default", session_server_instance_id=None, router_timeout=5.0, + tito_context_window=None, ) @@ -99,14 +103,17 @@ def _render(self, body: dict) -> httpx.Response: # Deliberately leak `stream: true` + stream_options the way a real # render would if the chat request streamed (vLLM copies the flag) — # the gateway must overwrite both before calling generate. - return httpx.Response(200, json={ - "request_id": "chatcmpl-render-1", - "token_ids": list(self.rendered_ids), - "sampling_params": sampling_params, - "model": body.get("model"), - "stream": True, - "stream_options": {"include_usage": True}, - }) + return httpx.Response( + 200, + json={ + "request_id": "chatcmpl-render-1", + "token_ids": list(self.rendered_ids), + "sampling_params": sampling_params, + "model": body.get("model"), + "stream": True, + "stream_options": {"include_usage": True}, + }, + ) def _generate(self, body: dict) -> httpx.Response: if self.generate_raw is not None: @@ -115,68 +122,73 @@ def _generate(self, body: dict) -> httpx.Response: # Faithful to v0.24.0: no logprobs block when sampling logprobs is null. logprobs = None if body.get("sampling_params", {}).get("logprobs") is not None: - logprobs = {"content": [ - {"token": f"token_id:{t}", "logprob": -0.1, "bytes": None, "top_logprobs": []} for t in ids - ]} - return httpx.Response(200, json={ - # v0.24.0 shape: no usage/model/created, random request_id. - "request_id": "9f0e6d1c", - "choices": [{ - "index": 0, - "finish_reason": self.finish_reason, - "token_ids": ids, - "logprobs": logprobs, - }], - "prompt_logprobs": None, - }) + logprobs = { + "content": [{"token": f"token_id:{t}", "logprob": -0.1, "bytes": None, "top_logprobs": []} for t in ids] + } + return httpx.Response( + 200, + json={ + # v0.24.0 shape: no usage/model/created, random request_id. + "request_id": "9f0e6d1c", + "choices": [ + { + "index": 0, + "finish_reason": self.finish_reason, + "token_ids": ids, + "logprobs": logprobs, + } + ], + "prompt_logprobs": None, + }, + ) def _derender(self, body: dict) -> httpx.Response: if self.derender_raw_content is not None: - return httpx.Response( - 200, content=self.derender_raw_content, headers={"content-type": "application/json"} - ) + return httpx.Response(200, content=self.derender_raw_content, headers={"content-type": "application/json"}) gen = body["generate_response"] prompt_tokens = body.get("prompt_tokens") or 0 completion_tokens = sum(len(c.get("token_ids") or []) for c in gen.get("choices", [])) - return httpx.Response(200, json={ - "id": gen.get("request_id", "x"), "object": "chat.completion", "created": 1, - "model": body["model"], - "choices": [{ - "index": c.get("index", 0), - # real derender passes finish_reason through verbatim — it - # never rewrites to "tool_calls"; the gateway does that. - "finish_reason": c.get("finish_reason"), - "message": dict(self.message), - "stop_reason": None, - } for c in gen.get("choices", [])], - "usage": { - "prompt_tokens": prompt_tokens, - "completion_tokens": completion_tokens, - "total_tokens": prompt_tokens + completion_tokens, + return httpx.Response( + 200, + json={ + "id": gen.get("request_id", "x"), + "object": "chat.completion", + "created": 1, + "model": body["model"], + "choices": [ + { + "index": c.get("index", 0), + # real derender passes finish_reason through verbatim — it + # never rewrites to "tool_calls"; the gateway does that. + "finish_reason": c.get("finish_reason"), + "message": dict(self.message), + "stop_reason": None, + } + for c in gen.get("choices", []) + ], + "usage": { + "prompt_tokens": prompt_tokens, + "completion_tokens": completion_tokens, + "total_tokens": prompt_tokens + completion_tokens, + }, }, - }) + ) @pytest.fixture() def server(tok, monkeypatch): - monkeypatch.setattr( - "agentix.tito.engine.session_app.load_tokenizer", lambda *a, **k: tok - ) + monkeypatch.setattr("agentix.tito.engine.session_app.load_tokenizer", lambda *a, **k: tok) pool = BackendPool([A]) srv = SessionServer(_args(), pool) replica = _VllmReplica() - srv._backend.client = httpx.AsyncClient( - transport=httpx.MockTransport(replica.handler), timeout=5.0 - ) + srv._backend.client = httpx.AsyncClient(transport=httpx.MockTransport(replica.handler), timeout=5.0) return srv, replica @pytest.fixture() def gateway(server): srv, replica = server - client = httpx.AsyncClient( - transport=httpx.ASGITransport(app=srv.app), base_url="http://gw", timeout=5.0 - ) + client = httpx.AsyncClient(transport=httpx.ASGITransport(app=srv.app), base_url="http://gw", timeout=5.0) return client, replica @@ -202,6 +214,37 @@ async def test_full_vllm_session_flow_over_http(gateway): assert got["metadata"]["accumulated_token_ids"] == prompt_ids + [7, 8] +@pytest.mark.asyncio +@pytest.mark.parametrize("field", ["content", "reasoning"]) +async def test_vllm_terminal_eos_is_not_rendered_twice(gateway, tok, field): + client, replica = gateway + replica.completion_ids = [7, tok.eos_token_id] + replica.message = {"role": "assistant", "content": None, field: "ok" + tok.eos_token} + sid = (await client.post("/sessions")).json()["session_id"] + response = await client.post(f"/sessions/{sid}/v1/chat/completions", json=_CHAT) + assert response.status_code == 200 + assert response.json()["choices"][0]["message"][field] == "ok" + recorded = (await client.get(f"/sessions/{sid}")).json() + assert ( + recorded["metadata"]["accumulated_token_ids"] + == replica.calls["generate"][0]["token_ids"] + replica.completion_ids + ) + assert response.json()["usage"]["completion_tokens"] == 2 + assert replica.calls["derender"][0]["generate_response"]["choices"][0]["token_ids"] == replica.completion_ids + + +@pytest.mark.asyncio +@pytest.mark.parametrize("finish,terminal", [("stop", False), ("length", True)]) +async def test_vllm_preserves_literal_eos_text_without_terminal_stop(gateway, tok, finish, terminal): + client, replica = gateway + replica.completion_ids = [7, tok.eos_token_id if terminal else 8] + replica.finish_reason = finish + replica.message = {"role": "assistant", "content": "ok" + tok.eos_token} + sid = (await client.post("/sessions")).json()["session_id"] + response = await client.post(f"/sessions/{sid}/v1/chat/completions", json=_CHAT) + assert response.json()["choices"][0]["message"]["content"] == "ok" + tok.eos_token + + @pytest.mark.asyncio async def test_vllm_second_turn_reuses_token_prefix(gateway): """The derendered assistant message echoed back with an appended tool turn @@ -244,7 +287,12 @@ async def test_vllm_forces_token_recording_fields(gateway): json={**_CHAT, "stream": True, "stream_options": {"include_usage": True}}, ) assert r.status_code == 200 - assert r.json()["choices"][0]["message"]["content"] == "ok done" + # The client asked for a stream, so the completed turn comes back as SSE. + assert r.headers["content-type"].startswith("text/event-stream") + lines = [line for line in r.text.split("\n") if line.startswith("data: ")] + events = [json.loads(line[6:]) for line in lines if line != "data: [DONE]"] + assert events[0]["choices"][0]["delta"]["content"] == "ok done" + assert events[-1]["choices"][0]["finish_reason"] == "stop" render_body = replica.calls["render"][0] assert render_body["logprobs"] is True @@ -268,9 +316,7 @@ async def test_vllm_explicit_null_top_logprobs_is_normalized(gateway): block, and every turn would 502. The gateway must coerce null to 0.""" client, replica = gateway sid = (await client.post("/sessions")).json()["session_id"] - r = await client.post( - f"/sessions/{sid}/v1/chat/completions", json={**_CHAT, "top_logprobs": None} - ) + r = await client.post(f"/sessions/{sid}/v1/chat/completions", json={**_CHAT, "top_logprobs": None}) assert r.status_code == 200 assert replica.calls["render"][0]["top_logprobs"] == 0 @@ -279,9 +325,7 @@ async def test_vllm_explicit_null_top_logprobs_is_normalized(gateway): async def test_vllm_agent_requested_top_logprobs_is_preserved(gateway): client, replica = gateway sid = (await client.post("/sessions")).json()["session_id"] - r = await client.post( - f"/sessions/{sid}/v1/chat/completions", json={**_CHAT, "top_logprobs": 5} - ) + r = await client.post(f"/sessions/{sid}/v1/chat/completions", json={**_CHAT, "top_logprobs": 5}) assert r.status_code == 200 assert replica.calls["render"][0]["top_logprobs"] == 5 @@ -316,10 +360,13 @@ async def test_vllm_tool_call_turn_rewrites_finish_reason(gateway): replica.message = { "role": "assistant", "content": None, - "tool_calls": [{ - "id": "call_1", "type": "function", - "function": {"name": "compute", "arguments": "{}"}, - }], + "tool_calls": [ + { + "id": "call_1", + "type": "function", + "function": {"name": "compute", "arguments": "{}"}, + } + ], } sid = (await client.post("/sessions")).json()["session_id"] r = await client.post(f"/sessions/{sid}/v1/chat/completions", json={**_CHAT, "tools": _TOOLS}) @@ -427,9 +474,7 @@ async def test_vllm_reasoning_is_mirrored_for_the_template_dialect(server): the stored trajectory message must carry the derendered reasoning under that key too, while the wire response keeps vLLM's `reasoning` verbatim.""" srv, replica = server - client = httpx.AsyncClient( - transport=httpx.ASGITransport(app=srv.app), base_url="http://gw", timeout=5.0 - ) + client = httpx.AsyncClient(transport=httpx.ASGITransport(app=srv.app), base_url="http://gw", timeout=5.0) replica.message = {"role": "assistant", "content": "ok done", "reasoning": "You are ok"} sid = (await client.post("/sessions")).json()["session_id"] r = await client.post(f"/sessions/{sid}/v1/chat/completions", json=_CHAT) @@ -466,7 +511,8 @@ async def test_vllm_generate_error_passes_through(gateway): async def test_vllm_derender_error_passes_through(gateway): client, replica = gateway replica.fail["derender"] = ( - 503, {"error": {"message": "parser overloaded", "type": "ServiceUnavailableError", "code": 503}} + 503, + {"error": {"message": "parser overloaded", "type": "ServiceUnavailableError", "code": 503}}, ) sid = (await client.post("/sessions")).json()["session_id"] r = await client.post(f"/sessions/{sid}/v1/chat/completions", json=_CHAT) @@ -483,9 +529,9 @@ async def test_vllm_malformed_generate_body_is_502(gateway): for weird in ( {"weird": True}, {"choices": []}, - {"choices": [{"index": 0, "finish_reason": "stop"}]}, # token_ids missing - {"choices": [{"index": 0, "token_ids": []}]}, # empty - {"choices": [{"index": 0, "token_ids": ["a", "b"]}]}, # non-int + {"choices": [{"index": 0, "finish_reason": "stop"}]}, # token_ids missing + {"choices": [{"index": 0, "token_ids": []}]}, # empty + {"choices": [{"index": 0, "token_ids": ["a", "b"]}]}, # non-int ): replica.generate_raw = weird r = await client.post(f"/sessions/{sid}/v1/chat/completions", json=_CHAT) @@ -503,9 +549,19 @@ async def test_vllm_logprobs_count_mismatch_is_502(gateway): sid = (await client.post("/sessions")).json()["session_id"] for weird in ( {"choices": [{"index": 0, "token_ids": [7, 8]}]}, # logprobs missing - {"choices": [{"index": 0, "token_ids": [7, 8], "logprobs": {"content": [ - {"token": "token_id:7", "logprob": -0.1, "bytes": None, "top_logprobs": []}, - ]}}]}, + { + "choices": [ + { + "index": 0, + "token_ids": [7, 8], + "logprobs": { + "content": [ + {"token": "token_id:7", "logprob": -0.1, "bytes": None, "top_logprobs": []}, + ] + }, + } + ] + }, ): replica.generate_raw = weird r = await client.post(f"/sessions/{sid}/v1/chat/completions", json=_CHAT) @@ -545,9 +601,7 @@ async def test_vllm_turn_record_retains_logprobs_and_render_skew(tok, monkeypatc client = httpx.AsyncClient(transport=httpx.ASGITransport(app=srv.app), base_url="http://gw", timeout=5.0) sid = (await client.post("/sessions")).json()["session_id"] - r = await client.post( - f"/sessions/{sid}/v1/chat/completions", json=_CHAT, headers={"x-request-id": "req-v1"} - ) + r = await client.post(f"/sessions/{sid}/v1/chat/completions", json=_CHAT, headers={"x-request-id": "req-v1"}) assert r.status_code == 200 lines = [json.loads(line) for line in (tmp_path / f"{sid}.jsonl").read_text().splitlines()] @@ -580,10 +634,13 @@ async def test_vllm_tool_call_record_carries_rewritten_finish_reason(tok, monkey replica.message = { "role": "assistant", "content": None, - "tool_calls": [{ - "id": "call_1", "type": "function", - "function": {"name": "compute", "arguments": "{}"}, - }], + "tool_calls": [ + { + "id": "call_1", + "type": "function", + "function": {"name": "compute", "arguments": "{}"}, + } + ], } srv._backend.client = httpx.AsyncClient(transport=httpx.MockTransport(replica.handler), timeout=5.0) client = httpx.AsyncClient(transport=httpx.ASGITransport(app=srv.app), base_url="http://gw", timeout=5.0) @@ -597,3 +654,50 @@ async def test_vllm_tool_call_record_carries_rewritten_finish_reason(tok, monkey (rec,) = [line for line in lines if line["schema_version"] == "tito.record.v1"] assert rec["finish_reason"] == "tool_calls" assert rec["assistant_message"]["tool_calls"][0]["id"] == "call_1" + + +@pytest.mark.asyncio +async def test_vllm_max_tokens_is_clamped_to_the_context_window(tok, monkeypatch): + """A per-turn budget larger than the room left after the session's exact + prompt ids is cut to that room; a budget that fits is forwarded verbatim. + Without the clamp vLLM rejects `prompt + max_tokens > max_model_len`, which + caps trajectory length at `max_model_len - max_tokens`.""" + monkeypatch.setattr("agentix.tito.engine.session_app.load_tokenizer", lambda *a, **k: tok) + + def make(window): + args = _args() + args.tito_context_window = window + srv = SessionServer(args, BackendPool([A])) + replica = _VllmReplica() + srv._backend.client = httpx.AsyncClient(transport=httpx.MockTransport(replica.handler), timeout=5.0) + client = httpx.AsyncClient(transport=httpx.ASGITransport(app=srv.app), base_url="http://gw", timeout=5.0) + return client, replica + + # Roomy window: the replica's render resolves max_tokens to 32 and it must + # pass through untouched. Also learn the gateway's exact prompt length. + client, replica = make(10_000) + sid = (await client.post("/sessions")).json()["session_id"] + assert (await client.post(f"/sessions/{sid}/v1/chat/completions", json=_CHAT)).status_code == 200 + generate_body = replica.calls["generate"][-1] + assert generate_body["sampling_params"]["max_tokens"] == 32 + prompt_len = len(generate_body["token_ids"]) + + # Window with only 5 tokens of room after that prompt (plus the render + # off-by-one margin): max_tokens -> 5, already in the chat request render + # sees (vLLM validates there too) and in the rendered sampling params. + from agentix.tito.engine.upstream import VllmUpstream + + client, replica = make(prompt_len + 5 + VllmUpstream.CONTEXT_MARGIN) + sid = (await client.post("/sessions")).json()["session_id"] + r = await client.post(f"/sessions/{sid}/v1/chat/completions", json={**_CHAT, "max_tokens": 4096}) + assert r.status_code == 200 + assert replica.calls["render"][-1]["max_tokens"] == 5 + generate_body = replica.calls["generate"][-1] + assert len(generate_body["token_ids"]) == prompt_len + assert generate_body["sampling_params"]["max_tokens"] == 5 + + # A prompt that already fills the window is left to the backend's verdict. + client, replica = make(prompt_len + VllmUpstream.CONTEXT_MARGIN) + sid = (await client.post("/sessions")).json()["session_id"] + assert (await client.post(f"/sessions/{sid}/v1/chat/completions", json=_CHAT)).status_code == 200 + assert replica.calls["generate"][-1]["sampling_params"]["max_tokens"] == 32 diff --git a/plugins/tito/tests/test_qwen3_5_golden.py b/plugins/tito/tests/test_qwen3_5_golden.py new file mode 100644 index 0000000..8ad5b55 --- /dev/null +++ b/plugins/tito/tests/test_qwen3_5_golden.py @@ -0,0 +1,195 @@ +"""Golden test: the incremental==from-scratch invariant on the REAL Qwen3.8 +tokenizer and its OWN chat template through the `qwen3_5` family. + +Qwen3.5 and Qwen3.8 share one template family (Qwen3.8 adds the +`reasoning_effort` system instruction and a `preserve_thinking` switch). The +family differs from Qwen3 in ways the tiny-vocab engine tests cannot see: +the template raises without a user turn, tool results render as a user turn +wrapped in ``, tool calls use the `` +XML dialect, and `` / `` are added tokens that must not split. +This module downloads the tokenizer-only files for Qwen/Qwen3.8-27B (a few +MB; the HF cache is reused) and drives a multi-turn tool-calling session +through `LinearTrajectory.prepare_prompt`, asserting the same invariants as +the Qwen3 golden test plus the `reasoning_effort` pin. + +Offline behavior: if the tokenizer is neither cached nor downloadable the +module SKIPS (marker: `network`) — it never fails a disconnected run. +""" + +from __future__ import annotations + +import json + +import pytest +from agentix.tito.engine.pretokenize import get_tito_tokenizer +from agentix.tito.engine.trajectory import LinearTrajectory, SessionRecord, SessionRegistry + +pytestmark = pytest.mark.network + +_REPO = "Qwen/Qwen3.8-27B" + + +@pytest.fixture(scope="module") +def qwen38_tok(): + from transformers import AutoTokenizer + + try: + return AutoTokenizer.from_pretrained(_REPO, local_files_only=True) + except Exception: + pass + try: + return AutoTokenizer.from_pretrained(_REPO) + except Exception as exc: # noqa: BLE001 - hub errors vary by transport + pytest.skip(f"Qwen3.8 tokenizer unavailable (offline?): {type(exc).__name__}: {exc}") + + +_TOOLS = [ + { + "type": "function", + "function": { + "name": "bash", + "description": "Run a shell command in the workspace.", + "parameters": { + "type": "object", + "properties": {"command": {"type": "string"}}, + "required": ["command"], + }, + }, + } +] + + +def _simulate_completion(tt, request_messages, assistant_message, prompt_ids, tools): + """The completion ids a template-canonical model emits: the from-scratch + render of request+assistant minus the prompt prefix, without the trailing + newline (the model stops at `<|im_end|>`).""" + full = tt.render_messages( + request_messages + [assistant_message], tools=tools, add_generation_prompt=False, tokenize=True + ) + assert full[: len(prompt_ids)] == prompt_ids, "assistant render must extend the generation prompt" + completion = full[len(prompt_ids):] + newline_id = tt.tokenizer.encode("\n", add_special_tokens=False)[0] + assert completion and completion[-1] == newline_id + return completion[:-1] + + +def _conversation(): + system = {"role": "system", "content": "You optimize GPU kernels. Be terse."} + user1 = {"role": "user", "content": "Make check.py pass, then make benchmark.py faster."} + assistant1 = { + "role": "assistant", + "content": "", + "reasoning_content": "First look at the task files before editing anything.", + "tool_calls": [ + { + "id": "call_0001", + "type": "function", + "function": {"name": "bash", "arguments": json.dumps({"command": "cat definition.json"})}, + } + ], + } + tool1 = {"role": "tool", "content": '{"op": "rmsnorm", "shape": [4096, 512]}', "tool_call_id": "call_0001"} + assistant2 = { + "role": "assistant", + "content": "", + "reasoning_content": "RMSNorm over the last dim; run the checker on the skeleton first.", + "tool_calls": [ + { + "id": "call_0002", + "type": "function", + "function": {"name": "bash", "arguments": json.dumps({"command": "python check.py"})}, + }, + { + "id": "call_0003", + "type": "function", + "function": {"name": "bash", "arguments": json.dumps({"command": "nvidia-smi -L"})}, + }, + ], + } + tool2a = {"role": "tool", "content": "FAILED: 8/8 workloads mismatch", "tool_call_id": "call_0002"} + tool2b = {"role": "tool", "content": "GPU 0: NVIDIA H100 80GB HBM3", "tool_call_id": "call_0003"} + assistant3 = { + "role": "assistant", + "content": "The skeleton fails all workloads; I will implement the kernel next.", + "reasoning_content": "Two tool results consumed; summarise and continue.", + } + turns = [ + ([system, user1], assistant1), + ([system, user1, assistant1, tool1], assistant2), + ([system, user1, assistant1, tool1, assistant2, tool2a, tool2b], assistant3), + ] + return turns + + +@pytest.mark.parametrize("effort", ["xhigh", "low"]) +def test_qwen3_8_incremental_equals_from_scratch_multi_turn_tool_calls(qwen38_tok, effort): + tt = get_tito_tokenizer( + qwen38_tok, "qwen3_5", allowed_append_roles=("tool",), chat_template_kwargs={"reasoning_effort": effort} + ) + registry = SessionRegistry(None, qwen38_tok, tito_tokenizer=tt) + tr = LinearTrajectory() + turns = _conversation() + + sources: list[list[str]] = [] + for request_messages, assistant in turns: + prepared = tr.prepare_prompt(request_messages, _TOOLS, tito_tokenizer=tt) + from_scratch = tt.render_messages(request_messages, tools=_TOOLS, add_generation_prompt=True, tokenize=True) + assert prepared.token_ids == from_scratch + assert prepared.prefix_stable is True + assert prepared.segments[0]["start"] == 0 + assert prepared.segments[-1]["end"] == len(prepared.token_ids) + for left, right in zip(prepared.segments, prepared.segments[1:], strict=False): + assert left["end"] == right["start"] + sources.append([s["source"] for s in prepared.segments]) + + completion = _simulate_completion(tt, request_messages, assistant, prepared.token_ids, _TOOLS) + tr.update_pretokenized_state( + request_messages, + assistant, + prompt_token_ids=prepared.token_ids, + completion_token_ids=completion, + max_trim_tokens=tt.max_trim_tokens, + ) + tr.append_record(SessionRecord( + timestamp=0.0, method="POST", path="/v1/chat/completions", status_code=200, + request={"model": "m", "messages": request_messages, "tools": _TOOLS}, response={}, + )) + + assert sources == [ + ["render"], + ["prefix", "tool", "generation_prompt"], + ["prefix", "tool", "generation_prompt"], # two consecutive tool results = one segment + ] + + final_messages = turns[-1][0] + [turns[-1][1]] + assert tt.fix_prefix(tr.token_ids) == tt.render_messages( + final_messages, tools=_TOOLS, add_generation_prompt=False, tokenize=True + ) + assert registry.compute_session_mismatch(tr) == [] + + # The pinned reasoning effort is in the rendered system prompt the model saw. + rendered = qwen38_tok.decode(tr.token_ids) + assert f"Reasoning effort is set to {effort}" in rendered + # Earlier-turn reasoning survives in the accumulated trajectory (preserve_thinking default). + assert "First look at the task files" in rendered + + +def test_qwen3_8_generation_prompt_opens_the_think_block(qwen38_tok): + tt = get_tito_tokenizer(qwen38_tok, "qwen3_5") + ids = tt.render_messages( + [{"role": "user", "content": "hi"}], tools=_TOOLS, add_generation_prompt=True, tokenize=True + ) + tail = qwen38_tok.decode(ids[-6:]) + assert tail.endswith("<|im_start|>assistant\n\n") + think_id = qwen38_tok.convert_tokens_to_ids("") + assert think_id in ids[-3:] # `` is a single added token, never split + + +def test_default_family_cannot_tokenize_qwen3_8_tool_results(qwen38_tok): + """Documents WHY the family exists: the model-agnostic engine's synthetic + context has no user turn and the Qwen3.5/3.8 template rejects it.""" + default = get_tito_tokenizer(qwen38_tok, "default") + turns = _conversation() + request1, assistant1 = turns[0] + with pytest.raises(ValueError): + default.tokenize_additional_non_assistant(request1 + [assistant1], turns[1][0], _TOOLS)