"""Stateful scrubber for reasoning/thinking blocks in streamed assistant text. The regex ``_strip_think_blocks`` is correct for a complete string but, run per-delta, erases an opening ```` that arrives alone, so downstream state machines leak reasoning. This class holds partial tags at delta boundaries until resolved; ``flush()`` releases held-back prose that was not a tag; ``reset()`` at the top of each turn. An open tag only starts a block at a block boundary (stream start / after a newline / whitespace-only line so far), so prose that *mentions* ```` is not suppressed; closed pairs are always suppressed (intentional). """ from __future__ import annotations import re from typing import Tuple __all__ = ["StreamingThinkScrubber", "THINK_TAG_NAMES", "THINK_OPEN_TAGS", "THINK_CLOSE_TAGS"] # The one list of model reasoning tag names. Every surface that hides reasoning (this scrubber, # the CLI stream filter, the gateway stream filter, the final-response regex stripper) binds to # these; a tag added here is covered everywhere. Consumers match case-insensitively, so the # literal tags are lowercase. THINK_TAG_NAMES: Tuple[str, ...] = ("think", "thinking", "reasoning", "thought", "REASONING_SCRATCHPAD") THINK_OPEN_TAGS: Tuple[str, ...] = tuple(f"<{name.lower()}>" for name in THINK_TAG_NAMES) THINK_CLOSE_TAGS: Tuple[str, ...] = tuple(f"" for name in THINK_TAG_NAMES) class StreamingThinkScrubber: """Stateful scrubber for streaming reasoning/thinking blocks. State: ``_in_block`` (inside an open block; text discarded), ``_buf`` (held-back partial-tag tail), ``_last_emitted_ended_newline`` (True iff the last emission ended with ``\\n`` or nothing was emitted yet — decides whether an open tag at buffer position 0 sits at a block boundary). """ # Literal tags so the hot path does string ops, not regex per feed(). _OPEN_TAGS: Tuple[str, ...] = THINK_OPEN_TAGS _CLOSE_TAGS: Tuple[str, ...] = THINK_CLOSE_TAGS _ALL_TAGS: Tuple[str, ...] = _OPEN_TAGS + _CLOSE_TAGS _MAX_TAG_LEN: int = max(len(tag) for tag in _ALL_TAGS) # Orphan close tag plus trailing whitespace (matches _strip_think_blocks case 3). _ORPHAN_CLOSE_RE = re.compile( "(?:" + "|".join(re.escape(t) for t in _CLOSE_TAGS) + r")[ \t\n\r]*", re.IGNORECASE ) def __init__(self) -> None: self.reset() def reset(self) -> None: """Reset all state. Call at the top of every new turn.""" self._in_block: bool = False self._buf: str = "" self._last_emitted_ended_newline: bool = True def _emit(self, out: list[str], text: str) -> None: """Append visible prose to *out* (orphan close tags stripped) and track the newline flag.""" text = self._strip_orphan_close_tags(text) if text: out.append(text) self._last_emitted_ended_newline = text.endswith("\n") def feed(self, text: str) -> str: """Feed one delta; return the scrubbed visible portion ("" when it is all reasoning or held back).""" if not text: return "" buf = self._buf + text self._buf = "" out: list[str] = [] while buf: if self._in_block: close_idx, close_len = self._find_first_tag(buf, self._CLOSE_TAGS) if close_idx == -1: # No close yet: hold back a possible partial close-tag prefix, drop the rest. self._hold_partial(buf, self._CLOSE_TAGS) break buf = buf[close_idx + close_len:] self._in_block = False continue # Priority 1: closed X pair anywhere (even inline pairs are almost # certainly leaked reasoning). Priority 2: unterminated open tag at a block # boundary (gated so prose mentioning '' isn't over-stripped). Earliest wins. pair = self._find_earliest_closed_pair(buf) open_idx, open_len = self._find_open_at_boundary(buf, out) if pair is not None and (open_idx == -1 or pair[0] <= open_idx): self._emit(out, buf[:pair[0]]) buf = buf[pair[1]:] continue if open_idx != -1: self._emit(out, buf[:open_idx]) self._in_block = True buf = buf[open_idx + open_len:] continue # No resolvable tag: hold back any partial-tag prefix at the tail # so a tag split across deltas isn't missed, then emit the rest. self._emit(out, self._hold_partial(buf, self._ALL_TAGS)) break return "".join(out) def _hold_partial(self, buf: str, tags: Tuple[str, ...]) -> str: """Move a trailing partial-tag prefix of *buf* into ``_buf``; return the remainder.""" held = self._max_partial_suffix(buf, tags) self._buf = buf[-held:] if held else "" return buf[:-held] if held else buf def flush(self) -> str: """End-of-stream flush: inside an unterminated block the held-back content is discarded (leaking partial reasoning is worse than a truncated answer), otherwise the tail is emitted verbatim. Always resets the boundary flag — intra-turn retries flush then stream again without ``reset()``, and a stale False flag made the new stream's opening ```` look mid-line.""" tail = "" if self._in_block else self._buf self._buf = "" self._in_block = False self._last_emitted_ended_newline = True return self._strip_orphan_close_tags(tail) if tail else "" # ── internal helpers ─────────────────────────────────────────────── @staticmethod def _find_first_tag(buf: str, tags: Tuple[str, ...]) -> Tuple[int, int]: """Return (earliest_index, tag_length) over *tags* (case-insensitive), or (-1, 0).""" buf_lower = buf.lower() hits = [(idx, len(tag)) for tag in tags if (idx := buf_lower.find(tag)) != -1] return min(hits) if hits else (-1, 0) def _find_earliest_closed_pair(self, buf: str): """(start_idx, end_idx) of the earliest ``...`` pair (non-greedy, case-insensitive), else None.""" buf_lower = buf.lower() pairs = [] for open_tag, close_tag in zip(self._OPEN_TAGS, self._CLOSE_TAGS): open_idx = buf_lower.find(open_tag) close_idx = buf_lower.find(close_tag, open_idx + len(open_tag)) if open_idx != -1 else -1 if close_idx != -1: pairs.append((open_idx, close_idx + len(close_tag))) return min(pairs) if pairs else None def _find_open_at_boundary(self, buf: str, already_emitted: list[str]) -> Tuple[int, int]: """Return the earliest block-boundary open-tag (idx, len), or (-1, 0).""" buf_lower = buf.lower() hits = [] for tag in self._OPEN_TAGS: idx = buf_lower.find(tag) while idx != -1 and not self._is_block_boundary(buf, idx, already_emitted): idx = buf_lower.find(tag, idx + 1) if idx != -1: hits.append((idx, len(tag))) return min(hits) if hits else (-1, 0) def _is_block_boundary(self, buf: str, idx: int, already_emitted: list[str]) -> bool: """True iff *idx* is a block boundary: position 0 after a newline-terminated (or no) prior emission, or any position whose preceding text on the current line is whitespace-only (when no newline precedes it in *buf*, the prior emission must also have ended with a newline).""" prior_newline = already_emitted[-1].endswith("\n") if already_emitted else self._last_emitted_ended_newline if idx == 0: return prior_newline preceding = buf[:idx] last_nl = preceding.rfind("\n") return (prior_newline if last_nl == -1 else True) and preceding[last_nl + 1:].strip() == "" @classmethod def _max_partial_suffix(cls, buf: str, tags: Tuple[str, ...]) -> int: """Longest buf-suffix that is a strict prefix of any tag (full matches are real tags, handled elsewhere).""" buf_lower = buf.lower() for i in range(min(len(buf_lower), cls._MAX_TAG_LEN - 1), 0, -1): suffix = buf_lower[-i:] if any(len(tag) > i and tag.startswith(suffix) for tag in tags): return i return 0 @classmethod def _strip_orphan_close_tags(cls, text: str) -> str: """Remove close tags with no matching open (always noise) plus trailing whitespace.""" return cls._ORPHAN_CLOSE_RE.sub("", text) if "