mirror of
https://github.com/turnstonelabs/turnstone.git
synced 2026-08-12 23:12:23 -06:00
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:
@@ -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")],
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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":
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user