[object Object]

← back to Exo

Drain tokens silently skipped in thinking parsing (#1898)

8cdc833892f019c4e120a924757bc461737673d3 · 2026-04-15 15:23:07 +0100 · rltakashige

## Motivation
Closes #1882

Files touched

Diff

commit 8cdc833892f019c4e120a924757bc461737673d3
Author: rltakashige <rl.takashige@gmail.com>
Date:   Wed Apr 15 15:23:07 2026 +0100

    Drain tokens silently skipped in thinking parsing   (#1898)
    
    ## Motivation
    Closes #1882
---
 .../runner/llm_inference/model_output_parsers.py   |  11 ++
 .../test_runner/test_finish_reason_sse.py          | 165 +++++++++++++++++++++
 2 files changed, 176 insertions(+)

diff --git a/src/exo/worker/runner/llm_inference/model_output_parsers.py b/src/exo/worker/runner/llm_inference/model_output_parsers.py
index c7e50649..a5cf19e6 100644
--- a/src/exo/worker/runner/llm_inference/model_output_parsers.py
+++ b/src/exo/worker/runner/llm_inference/model_output_parsers.py
@@ -330,6 +330,12 @@ def parse_thinking_models(
     """
     is_thinking = starts_in_thinking
     accumulated = ""
+    pending_buffer: list[GenerationResponse] = []
+
+    def drain_pending(_is_thinking: bool):
+        for buffered in pending_buffer:
+            yield buffered.model_copy(update={"is_thinking": _is_thinking})
+        pending_buffer.clear()
 
     for response in responses:
         if response is None:
@@ -339,25 +345,30 @@ def parse_thinking_models(
         accumulated += response.text
 
         if response.finish_reason is not None:
+            yield from drain_pending(is_thinking)
             yield response.model_copy(update={"is_thinking": False})
             continue
 
         if accumulated == think_start and not is_thinking:
             is_thinking = True
             accumulated = ""
+            pending_buffer.clear()
             continue
         if accumulated == think_end and is_thinking:
             is_thinking = False
             accumulated = ""
+            pending_buffer.clear()
             continue
 
         if (think_start and accumulated == think_start[: len(accumulated)]) or (
             think_end and accumulated == think_end[: len(accumulated)]
         ):
+            pending_buffer.append(response)
             continue
 
         accumulated = ""
 
+        yield from drain_pending(is_thinking)
         yield response.model_copy(update={"is_thinking": is_thinking})
 
 
diff --git a/src/exo/worker/tests/unittests/test_runner/test_finish_reason_sse.py b/src/exo/worker/tests/unittests/test_runner/test_finish_reason_sse.py
index 0aee03ba..90afe7f4 100644
--- a/src/exo/worker/tests/unittests/test_runner/test_finish_reason_sse.py
+++ b/src/exo/worker/tests/unittests/test_runner/test_finish_reason_sse.py
@@ -1,3 +1,4 @@
+import json
 from collections.abc import Generator
 from typing import Any
 
@@ -403,3 +404,167 @@ class TestBatchGeneratorSingleNext:
         assert _got_finish(collected), (
             f"No finish_reason in collected: {[(type(r).__name__, getattr(r, 'finish_reason', None) if isinstance(r, GenerationResponse) else 'tool') for r in collected]}"
         )
+
+
+# ── parse_thinking_models prefix buffering ──────────────────────
+
+
+def _drain_text(
+    results: list[GenerationResponse | ToolCallResponse],
+) -> str:
+    return "".join(
+        r.text
+        for r in results
+        if isinstance(r, GenerationResponse) and r.finish_reason is None
+    )
+
+
+class TestThinkingModelsPrefixBuffering:
+    def test_lone_lt_is_preserved(self):
+        tokens = [
+            _make_response("<", 0),
+            _make_response("function", 1),
+            _make_response(">", 2),
+            _make_response("", 3, finish_reason="stop"),
+        ]
+        results = _step_until_finish(
+            parse_thinking_models(
+                _queue_source(tokens),
+                think_start="<think>",
+                think_end="</think>",
+                starts_in_thinking=False,
+            )
+        )
+        assert _drain_text(results) == "<function>"
+        gens = [r for r in results if isinstance(r, GenerationResponse)]
+        assert all(not r.is_thinking for r in gens)
+
+    def test_lone_lt_slash_is_preserved(self):
+        tokens = [
+            _make_response("</", 0),
+            _make_response("parameter", 1),
+            _make_response(">", 2),
+            _make_response("", 3, finish_reason="stop"),
+        ]
+        results = _step_until_finish(
+            parse_thinking_models(
+                _queue_source(tokens),
+                think_start="<think>",
+                think_end="</think>",
+                starts_in_thinking=False,
+            )
+        )
+        assert _drain_text(results) == "</parameter>"
+
+    def test_partial_prefix_then_diverge(self):
+        tokens = [
+            _make_response("<", 0),
+            _make_response("t", 1),
+            _make_response("h", 2),
+            _make_response("other", 3),
+            _make_response("", 4, finish_reason="stop"),
+        ]
+        results = _step_until_finish(
+            parse_thinking_models(
+                _queue_source(tokens),
+                think_start="<think>",
+                think_end="</think>",
+                starts_in_thinking=False,
+            )
+        )
+        assert _drain_text(results) == "<thother"
+
+    def test_real_think_tag_still_swallowed(self):
+        tokens = [
+            _make_response("<", 0),
+            _make_response("think", 1),
+            _make_response(">", 2),
+            _make_response("body", 3),
+            _make_response("</", 4),
+            _make_response("think", 5),
+            _make_response(">", 6),
+            _make_response("after", 7),
+            _make_response("", 8, finish_reason="stop"),
+        ]
+        results = _step_until_finish(
+            parse_thinking_models(
+                _queue_source(tokens),
+                think_start="<think>",
+                think_end="</think>",
+                starts_in_thinking=False,
+            )
+        )
+        gens = [
+            r
+            for r in results
+            if isinstance(r, GenerationResponse) and r.finish_reason is None
+        ]
+        texts = [(r.text, r.is_thinking) for r in gens]
+        assert texts == [("body", True), ("after", False)]
+
+    def test_finish_reason_flushes_buffer(self):
+        tokens = [
+            _make_response("<", 0),
+            _make_response("", 1, finish_reason="stop"),
+        ]
+        results = _step_until_finish(
+            parse_thinking_models(
+                _queue_source(tokens),
+                think_start="<think>",
+                think_end="</think>",
+                starts_in_thinking=False,
+            )
+        )
+        gens = [r for r in results if isinstance(r, GenerationResponse)]
+        assert len(gens) == 2
+        assert gens[0].text == "<"
+        assert gens[0].is_thinking is False
+        assert gens[0].finish_reason is None
+        assert gens[1].finish_reason == "stop"
+        assert gens[1].is_thinking is False
+
+    def test_tool_call_after_prefix_tokens_parses(self):
+        def _capture_parser(text: str) -> dict[str, Any]:
+            return {"name": "captured", "arguments": {"raw": text}}
+
+        tool_parser = make_mlx_parser("<tool_call>", "</tool_call>", _capture_parser)
+
+        tokens = [
+            _make_response("<tool_call>", 0),
+            _make_response("\n", 1),
+            _make_response("<", 2),
+            _make_response("function", 3),
+            _make_response("=glob", 4),
+            _make_response(">", 5),
+            _make_response("\n", 6),
+            _make_response("<", 7),
+            _make_response("parameter", 8),
+            _make_response("=pattern", 9),
+            _make_response(">", 10),
+            _make_response("**/*", 11),
+            _make_response("</", 12),
+            _make_response("parameter", 13),
+            _make_response(">", 14),
+            _make_response("</", 15),
+            _make_response("function", 16),
+            _make_response(">", 17),
+            _make_response("</tool_call>", 18, finish_reason="stop"),
+        ]
+
+        thinking = parse_thinking_models(
+            _queue_source(tokens),
+            think_start="<think>",
+            think_end="</think>",
+            starts_in_thinking=False,
+        )
+        results = _step_until_finish(
+            parse_tool_calls(thinking, tool_parser, tools=None)
+        )
+
+        tool_results = [r for r in results if isinstance(r, ToolCallResponse)]
+        assert len(tool_results) == 1
+        raw = json.loads(tool_results[0].tool_calls[0].arguments)["raw"]  # pyright: ignore[reportAny]
+        assert "<function=glob>" in raw
+        assert "<parameter=pattern>" in raw
+        assert "</parameter>" in raw
+        assert "</function>" in raw

← 2cd66ae4 Fix out of order event idx causing fatal crashes (#1894)  ·  back to Exo  ·  Add usage stats to tool calls and handle multiple tool calls 87329c80 →