refactor(session): extract think-tag splitting into ThinkTagSplitter

The interactive chunk consumer's _flush_text/_drain_pending closure pair
carried the partial-tag carry buffer and in-think state inline. The
tag-scanning half moves to turnstone/core/streaming_text.py as a
standalone ThinkTagSplitter (carry buffer, in_think state, earliest-
index tag selection, MAX_TAG_LEN safe-flush); dispatch and accumulation
stay in the session behind the emit callback, and out-of-band
transitions (reasoning_delta path, tool-call starts, cancellation)
read/write splitter.in_think and flush_pending() where they previously
touched the closure locals.

Pure move: table-driven pins covering partial-tag buffering across
chunk boundaries, the safe-flush margin, open/close tag precedence,
in_think transitions, and reasoning-vs-content dispatch were written
against the closure implementation and pass unchanged against the
extracted class — byte-identical emitted text, identical UI callback
ordering. The session-level _THINK_*/_MAX_TAG_LEN class constants fold
into the class.
This commit is contained in:
Patrick Buckley
2026-08-03 23:48:30 -07:00
parent a1dfe0bd4f
commit 3b9de67e8c
5 changed files with 296 additions and 72 deletions
+3 -2
View File
@@ -418,8 +418,9 @@ class TestStreamFlushBeforeToolCalls:
arguments_delta: str = ""
def stream_content_then_tool():
# Content long enough to leave chars in pending buffer
# (_MAX_TAG_LEN = 13, so _drain_pending retains last 13 chars)
# Content long enough to leave chars in the tag-scan carry
# buffer (ThinkTagSplitter retains the last MAX_TAG_LEN = 12
# chars until a flush)
yield FakeChunk(content_delta="Hello world, this is a test message")
yield FakeChunk(
tool_call_deltas=[FakeToolDelta(index=0, id="tc_1", name="bash")],
+3 -2
View File
@@ -25,6 +25,7 @@ from turnstone.core.memory import load_last_error
from turnstone.core.providers import StreamChunk, UsageInfo
from turnstone.core.providers._protocol import IncompleteStreamError
from turnstone.core.session import ChatSession
from turnstone.core.streaming_text import ThinkTagSplitter
from turnstone.core.trajectory import dicts_from_turns
@@ -175,10 +176,10 @@ class TestMidStreamRetry:
assert len(assistant) == 1
assert assistant[0]["content"] == "Hello world"
# The dead attempt's tokens streamed live (minus the trailing
# _MAX_TAG_LEN chars the tag-scan buffer held back — never
# MAX_TAG_LEN chars the tag-scan buffer held back — never
# displayed, so nothing needs to finalize them); the retried text
# is a fresh bubble, not an append onto the dead attempt's.
flushed = first_text[: len(first_text) - ChatSession._MAX_TAG_LEN]
flushed = first_text[: len(first_text) - ThinkTagSplitter.MAX_TAG_LEN]
assert ui.of("content") == [flushed, "Hello world"]
# The dead attempt is finalized EVERYWHERE before re-streaming:
# stream_end (client finalize) -> turn_committed (server buffer
+196
View File
@@ -0,0 +1,196 @@
"""Behavior pins for the interactive think-tag splitting layer.
``ChatSession._stream_attempt`` splits streamed content into content vs
reasoning around ``<think>``/``<reasoning>`` tags, buffering potential
partial tags across chunk boundaries. These tables pin the CURRENT
emission behavior — exact UI token sequence and final message content —
so the logic can move into a standalone ``ThinkTagSplitter`` class with
byte-identical output. Every case drives the real chunk consumer end to
end; none reaches into the implementation, so the same rows must stay
green across the extraction.
Pinned rules:
- partial-tag buffering: a chunk ending in a possible tag prefix emits
nothing until the tag resolves or the safe-flush margin clears it;
- safe-flush: with no tag in sight, everything but the trailing
MAX-tag-length chars flushes immediately (live streaming), the tail
only at stream end;
- open/close tag selection by EARLIEST index among the tag variants;
- ``in_think`` transitions, including the ``reasoning_delta`` (path-1)
interplay and the raw pending flush when tool calls begin (pending
text can no longer be a partial tag).
"""
import pytest
from tests._session_helpers import make_session
from turnstone.core.providers import StreamChunk, ToolCallDelta
from turnstone.core.streaming_text import ThinkTagSplitter
class _TokenRecorderUI:
"""Records dispatched token events; stubs the rest of the surface
``_stream_attempt`` touches."""
def __init__(self):
self.tokens = []
def on_content_token(self, text):
self.tokens.append(("content", text))
def on_reasoning_token(self, text):
self.tokens.append(("reasoning", text))
def on_thinking_stop(self):
pass
def on_stream_end(self):
pass
def on_info(self, message):
pass
def on_error(self, message):
pass
def _drive(chunks, *, show_reasoning=True):
session = make_session()
session.show_reasoning = show_reasoning
ui = _TokenRecorderUI()
session.ui = ui
msg = session._stream_attempt(iter(chunks))
return msg, ui.tokens
def _c(text):
return StreamChunk(content_delta=text)
_FINISH = StreamChunk(finish_reason="stop")
# (case id, chunks, expected ordered token events, expected final content)
CASES = [
(
"plain_short_held_until_end",
[_c("hello"), _FINISH],
[("content", "hello")],
"hello",
),
(
"think_block_single_chunk",
[_c("<think>deep</think>answer"), _FINISH],
[("reasoning", "deep"), ("content", "answer")],
"answer",
),
(
"tag_split_across_chunk_boundary",
[_c("before<thi"), _c("nk>r</think>after"), _FINISH],
[("content", "before"), ("reasoning", "r"), ("content", "after")],
"beforeafter",
),
(
"reasoning_tag_variant",
[_c("<reasoning>x</reasoning>done"), _FINISH],
[("reasoning", "x"), ("content", "done")],
"done",
),
(
"earliest_tag_index_wins",
[_c("a<reasoning>b</reasoning>c<think>d</think>e"), _FINISH],
[
("content", "a"),
("reasoning", "b"),
("content", "c"),
("reasoning", "d"),
("content", "e"),
],
"ace",
),
(
"safe_flush_streams_all_but_max_tag_len",
[_c("x" * 30), _FINISH],
[("content", "x" * 18), ("content", "x" * 12)],
"x" * 30,
),
(
"path1_reasoning_then_content",
[StreamChunk(reasoning_delta="rr"), _c("cc"), _FINISH],
[("reasoning", "rr"), ("content", "cc")],
"cc",
),
(
"unterminated_think_tail_flushes_as_reasoning",
[_c("<think>abc"), _FINISH],
[("reasoning", "abc")],
"",
),
(
"content_surrounds_think_block",
[_c("before<think>mid</think>after"), _FINISH],
[("content", "before"), ("reasoning", "mid"), ("content", "after")],
"beforeafter",
),
(
"safe_flush_inside_think_block",
[_c("<think>" + "y" * 20), _c("</think>ok"), _FINISH],
[("reasoning", "y" * 8), ("reasoning", "y" * 12), ("content", "ok")],
"ok",
),
]
@pytest.mark.parametrize(
("chunks", "expected_events", "expected_content"),
[c[1:] for c in CASES],
ids=[c[0] for c in CASES],
)
def test_tag_splitting_emissions(chunks, expected_events, expected_content):
msg, tokens = _drive(chunks)
assert tokens == expected_events
assert msg["content"] == expected_content
def test_show_reasoning_off_suppresses_reasoning_dispatch_only():
msg, tokens = _drive([_c("<think>deep</think>answer"), _FINISH], show_reasoning=False)
assert tokens == [("content", "answer")]
assert msg["content"] == "answer"
def test_splitter_standalone_contract():
"""The extracted class is consumable without a session: feed() emits
tag-resolved spans through the callback, flush_pending() drains the
carry buffer raw, and in_think is externally writable for
out-of-band transitions."""
events = []
splitter = ThinkTagSplitter(lambda text, is_reasoning: events.append((text, is_reasoning)))
splitter.feed("a<think>b</think>c")
assert events == [("a", False), ("b", True)]
assert splitter.pending == "c"
splitter.flush_pending()
assert events == [("a", False), ("b", True), ("c", False)]
assert splitter.pending == ""
splitter.in_think = True
splitter.feed("tail")
splitter.flush_pending()
assert events[-1] == ("tail", True)
def test_tool_calls_flush_pending_raw_at_current_state():
# Once tool calls begin, buffered text cannot be a partial tag: it
# flushes RAW (no tag scan) at the current in_think state.
chunks = [
_c("part<thi"),
StreamChunk(tool_call_deltas=[ToolCallDelta(index=0, id="tc1", name="bash")]),
StreamChunk(
tool_call_deltas=[ToolCallDelta(index=0, arguments_delta="{}")],
finish_reason="tool_calls",
),
]
msg, tokens = _drive(chunks)
assert tokens == [("content", "part<thi")]
assert msg["content"] == "part<thi"
assert msg["tool_calls"] == [
{"id": "tc1", "type": "function", "function": {"name": "bash", "arguments": "{}"}}
]
+12 -68
View File
@@ -190,6 +190,7 @@ from turnstone.core.storage._utils import (
attachment_to_content_part,
normalize_search_terms,
)
from turnstone.core.streaming_text import ThinkTagSplitter
from turnstone.core.tool_advisory import (
make_system_turn,
render_output_guard_text,
@@ -7748,12 +7749,6 @@ class ChatSession:
log.warning("stream.retry.recreate_failed", exc_info=True)
raise e from None
# Tags that delimit reasoning blocks in content stream.
# Checked in order; first match wins.
_THINK_OPEN_TAGS = ("<think>", "<reasoning>")
_THINK_CLOSE_TAGS = ("</think>", "</reasoning>")
_MAX_TAG_LEN = max(len(t) for t in _THINK_OPEN_TAGS + _THINK_CLOSE_TAGS)
def _stream_attempt(
self, stream: Iterator[StreamChunk], my_generation: int = 0
) -> dict[str, Any]:
@@ -7789,9 +7784,7 @@ class ChatSession:
provider_blocks: list[dict[str, Any]] = []
usage_acc: UsageInfo | None = None
first_token = True
in_think = False # inside a <think>...</think> block
path1_reasoning = False # last reasoning came via reasoning_delta field
pending = "" # buffer for partial tag detection
def _flush_text(text: str, is_reasoning: bool) -> None:
"""Dispatch text to the appropriate UI callback."""
@@ -7805,53 +7798,9 @@ class ChatSession:
content_parts.append(text)
self.ui.on_content_token(text)
def _drain_pending() -> None:
"""Process the pending buffer, flushing content and detecting tags."""
nonlocal pending, in_think
while pending:
if in_think:
# Look for any close tag
best_idx, best_tag = None, None
for tag in self._THINK_CLOSE_TAGS:
idx = pending.find(tag)
if idx != -1 and (best_idx is None or idx < best_idx):
best_idx, best_tag = idx, tag
if best_idx is not None:
assert best_tag is not None
_flush_text(pending[:best_idx], True)
pending = pending[best_idx + len(best_tag) :]
in_think = False
continue
# No close tag found — check if tail could be a partial tag
safe = len(pending) - self._MAX_TAG_LEN
if safe > 0:
_flush_text(pending[:safe], True)
pending = pending[safe:]
break
else:
# Look for any open tag
best_idx, best_tag = None, None
for tag in self._THINK_OPEN_TAGS:
idx = pending.find(tag)
if idx != -1 and (best_idx is None or idx < best_idx):
best_idx, best_tag = idx, tag
if best_idx is not None:
assert best_tag is not None
_flush_text(pending[:best_idx], False)
pending = pending[best_idx + len(best_tag) :]
in_think = True
continue
# No open tag found — flush all but potential partial tag
safe = len(pending) - self._MAX_TAG_LEN
if safe > 0:
_flush_text(pending[:safe], False)
pending = pending[safe:]
break
# Owns the partial-tag carry buffer and the in-think state;
# dispatch stays here in _flush_text.
splitter = ThinkTagSplitter(_flush_text)
def _stop_spinner_once() -> None:
"""Stop the spinner on first real content. Call is idempotent."""
@@ -7878,8 +7827,7 @@ class ChatSession:
appends to ``content``, hiding the partial-output signal
from the model.
"""
if pending:
_flush_text(pending, in_think)
splitter.flush_pending()
self.ui.on_stream_end()
partial: dict[str, Any] = {"role": "assistant"}
partial["content"] = "".join(content_parts) or ""
@@ -7930,7 +7878,7 @@ class ChatSession:
if chunk.reasoning_delta:
_stop_spinner_once()
reasoning_parts.append(chunk.reasoning_delta)
in_think = True
splitter.in_think = True
path1_reasoning = True
if self.show_reasoning:
self.ui.on_reasoning_token(chunk.reasoning_delta)
@@ -7941,21 +7889,18 @@ class ChatSession:
# Close reasoning if transitioning from Path 1 reasoning
if path1_reasoning:
path1_reasoning = False
in_think = False
pending += chunk.content_delta
_drain_pending()
splitter.in_think = False
splitter.feed(chunk.content_delta)
# Handle tool call deltas
if chunk.tool_call_deltas:
_stop_spinner_once()
# Flush any buffered content — model has moved to tool calls,
# so pending text cannot be a partial <think> tag.
if pending:
_flush_text(pending, in_think)
pending = ""
splitter.flush_pending()
# Close reasoning if transitioning from reasoning
if in_think:
in_think = False
if splitter.in_think:
splitter.in_think = False
for tcd in chunk.tool_call_deltas:
# THE tool-call merge rule, shared with drain_stream
# and the Google raw-fidelity capture — the chat
@@ -7985,8 +7930,7 @@ class ChatSession:
raise
# Flush any remaining buffered text
if pending:
_flush_text(pending, in_think)
splitter.flush_pending()
# Warn on non-standard finish reasons
if finish_reason == "length":
+82
View File
@@ -0,0 +1,82 @@
"""Streaming think-tag splitting for interactive content streams."""
from __future__ import annotations
from typing import TYPE_CHECKING
if TYPE_CHECKING:
from collections.abc import Callable
class ThinkTagSplitter:
"""Split streamed content into content vs reasoning around think tags.
Local-model servers without a reasoning parser emit reasoning inline
as ``<think>``/``<reasoning>`` blocks inside the content stream, and
a tag can arrive split across chunk boundaries. This class owns the
carry buffer and the in-tag state: :meth:`feed` buffers a chunk and
emits every span that provably cannot be part of an unresolved tag;
the trailing :data:`MAX_TAG_LEN` chars are held until a later chunk
(or a final :meth:`flush_pending`) resolves them.
Emission goes through the *emit* callback ``(text, is_reasoning)`` —
spans arrive in stream order, never empty, exactly once. The
consumer owns dispatch (UI callbacks, accumulation) and may
read/write :attr:`in_think` directly to mirror out-of-band
transitions (the provider-normalized ``reasoning_delta`` path,
tool-call starts).
Tag selection: at each step the EARLIEST occurrence wins among the
tag variants for the current state (open tags outside a block, close
tags inside).
"""
OPEN_TAGS: tuple[str, ...] = ("<think>", "<reasoning>")
CLOSE_TAGS: tuple[str, ...] = ("</think>", "</reasoning>")
MAX_TAG_LEN = max(len(t) for t in OPEN_TAGS + CLOSE_TAGS)
def __init__(self, emit: Callable[[str, bool], None]) -> None:
self._emit = emit
self.pending = ""
self.in_think = False
def feed(self, text: str) -> None:
"""Buffer *text* and emit every tag-resolved span."""
self.pending += text
self._drain()
def flush_pending(self) -> None:
"""Emit the raw carry buffer at the current state, with no tag scan.
For stream boundaries where a partial tag is no longer possible:
end of stream, tool calls beginning, cancellation.
"""
if self.pending:
self._emit(self.pending, self.in_think)
self.pending = ""
def _drain(self) -> None:
while self.pending:
tags = self.CLOSE_TAGS if self.in_think else self.OPEN_TAGS
best_idx, best_tag = None, None
for tag in tags:
idx = self.pending.find(tag)
if idx != -1 and (best_idx is None or idx < best_idx):
best_idx, best_tag = idx, tag
if best_idx is not None:
assert best_tag is not None
if best_idx:
self._emit(self.pending[:best_idx], self.in_think)
self.pending = self.pending[best_idx + len(best_tag) :]
self.in_think = not self.in_think
continue
# No tag in sight — everything but a possible partial-tag tail
# is provably safe to emit now (live streaming beats holding
# the whole buffer for a tag that may never come).
safe = len(self.pending) - self.MAX_TAG_LEN
if safe > 0:
self._emit(self.pending[:safe], self.in_think)
self.pending = self.pending[safe:]
break