# SPDX-License-Identifier: Apache-2.0 """ Thinking/reasoning content parser for separating ... blocks. Provides both streaming (ThinkingParser) and non-streaming (extract_thinking) interfaces for separating reasoning content from regular response content. Used by reasoning models like DeepSeek R1, Qwen3/3.5, MiniMax that wrap their chain-of-thought reasoning in ... tags. """ import re from collections.abc import Callable, Sequence from typing import List, Optional, Tuple # Tags used for thinking blocks _OPEN_TAG = "" _CLOSE_TAG = "" _OPEN_LEN = len(_OPEN_TAG) # 7 _CLOSE_LEN = len(_CLOSE_TAG) # 8 _MINIMAX_OPEN_TAG = "" _MINIMAX_CLOSE_TAG = "" _HY3_OPEN_TAG = "" _HY3_CLOSE_TAG = "" # Regex for non-streaming extraction (complete text) _THINKING_PATTERN = re.compile(r'(.*?)', re.DOTALL) # Handle case where is missing but is present # (scheduler prepends \n but the tag may be split) _THINKING_TAIL_PATTERN = re.compile(r'^(.*?)', re.DOTALL) def _safe_tokenizer_attr(tokenizer, attr: str, default=None): if tokenizer is None: return default try: return getattr(tokenizer, attr, default) except (AttributeError, TypeError, ValueError): return default def _single_token_id(value) -> int | None: if value is None: return None try: return int(value) except (TypeError, ValueError): return None def _convert_token_to_id(tokenizer, token: str) -> int | None: convert = _safe_tokenizer_attr(tokenizer, "convert_tokens_to_ids") if not callable(convert): return None try: token_id = convert(token) except (AttributeError, KeyError, TypeError, ValueError): return None if token_id == _safe_tokenizer_attr(tokenizer, "unk_token_id"): return None return _single_token_id(token_id) def _encode_prompt_ids(tokenizer, prompt: str) -> list[int] | None: encode = _safe_tokenizer_attr(tokenizer, "encode") if not callable(encode): return None try: return list(encode(prompt, add_special_tokens=False)) except TypeError: try: return list(encode(prompt)) except Exception: return None except Exception: return None def _think_end_token_ids(tokenizer) -> list[int] | None: think_end_id = _single_token_id(_safe_tokenizer_attr(tokenizer, "think_end_id")) if think_end_id is not None: return [think_end_id] think_end_tag = _safe_tokenizer_attr(tokenizer, "think_end", _CLOSE_TAG) encoded = _encode_prompt_ids(tokenizer, think_end_tag or _CLOSE_TAG) if encoded: return encoded token_id = _convert_token_to_id(tokenizer, _CLOSE_TAG) if token_id is not None: return [token_id] return None def prompt_opens_thinking( tokenizer, prompt: str, prompt_token_ids: Sequence[int] | None = None, ) -> tuple[bool, str]: """Return whether a raw prompt would make the engine prepend ````. Presentation-layer stripping must mirror the engine/scheduler decision, not just the raw text suffix. Some prompts can contain a literal ```` without tokenizing to the model's think-start id, and templates can leave the think-start token in the final token tail without the raw string ending in the visible tag. When the caller already has prompt ids from the same tokenizer path as the scheduler, those ids are authoritative. """ think_tag = ( _safe_tokenizer_attr(tokenizer, "think_start", _OPEN_TAG) or _OPEN_TAG ) if tokenizer is None: return prompt.rstrip().endswith(think_tag), think_tag think_start_id = _single_token_id( _safe_tokenizer_attr(tokenizer, "think_start_id") ) if think_start_id is None: think_start_id = _convert_token_to_id(tokenizer, think_tag) if think_start_id is None: return False, think_tag if prompt_token_ids is None: prompt_ids = _encode_prompt_ids(tokenizer, prompt) else: prompt_ids = list(prompt_token_ids) if not prompt_ids or not think_start_id: return False, think_tag last_tokens = list(prompt_ids[-3:]) if think_start_id not in last_tokens: return False, think_tag last_idx = len(last_tokens) - 1 - last_tokens[::-1].index(think_start_id) after_start = last_tokens[last_idx + 1 :] if after_start: think_end_ids = _think_end_token_ids(tokenizer) if think_end_ids and think_end_ids[0] in after_start: return False, think_tag return True, think_tag def extract_thinking(text: str) -> Tuple[str, str]: """Extract thinking and content from complete text. Handles: - Normal: ``reasoninganswer`` → ``("reasoning", "answer")`` - No thinking: ``just answer`` → ``("", "just answer")`` - Partial (no open tag): ``reasoninganswer`` → ``("reasoning", "answer")`` - Empty think: ``answer`` → ``("", "answer")`` - Think only: ``reasoning`` → ``("reasoning", "")`` - Malformed (open with no close): ``everything…`` → ``("", "everything…")`` — recovery for V4-style models that occasionally skip the ```` boundary token. Without this fallback the entire body would be classified as thinking and the visible answer would be empty. Tag-free text is always classified as content. Mirrors ``ThinkingParser.finish()`` recovery semantics (`_content_emitted` fallback): when the model emits no thinking markers, surface the body as the answer so the response is never empty. Args: text: Complete model output text. Returns: Tuple of (thinking_content, regular_content). """ if not text: return ("", "") text = ( text.replace(_MINIMAX_OPEN_TAG, _OPEN_TAG) .replace(_MINIMAX_CLOSE_TAG, _CLOSE_TAG) .replace(_HY3_OPEN_TAG, _OPEN_TAG) .replace(_HY3_CLOSE_TAG, _CLOSE_TAG) ) thinking_parts = [] remaining = text # Extract all ... blocks while True: match = _THINKING_PATTERN.search(remaining) if not match: break thinking_parts.append(match.group(1)) remaining = remaining[:match.start()] + remaining[match.end():] if thinking_parts: thinking = "\n".join(thinking_parts).strip() return (thinking, remaining.strip()) # Handle partial: content before without tag if '' in text and '' not in text: match = _THINKING_TAIL_PATTERN.match(text) if match: thinking = match.group(1).strip() remaining = text[match.end():].strip() return (thinking, remaining) # Malformed: opened but never closed. Drop the open tag and # treat the remainder as content so the answer body is not empty. if '' in text and '' not in text: idx = text.index('') before = text[:idx] after = text[idx + _OPEN_LEN:] return ("", (before + after).strip()) return ("", text) class ThinkingParser: """Stateful streaming parser for separating ... from content. Handles streaming chunks where tags may span multiple chunks. Returns (thinking_delta, content_delta) tuples for each feed() call. Example:: parser = ThinkingParser() # Chunk 1: "Let me" t, c = parser.feed("Let me") # t = "Let me", c = "" # Chunk 2: " thinkAnswer" t, c = parser.feed(" thinkAnswer") # t = " think", c = "Answer" # Flush remaining t, c = parser.finish() """ def __init__(self, start_in_thinking: bool = False): self._in_thinking: bool = start_in_thinking self._buffer: str = "" # Buffer for potential partial tags # Recovery state for malformed thinking: when the prompt prepends # ```` and the model never emits ```` before EOS, # everything we streamed went out as thinking. The streamed events # cannot be retracted, so finish() emits the accumulated thinking # text once more as content — the client will show both panels but # the answer body is no longer empty. self._close_seen: bool = False self._thinking_accumulated: List[str] = [] self._content_emitted: bool = False def feed(self, text: str) -> Tuple[str, str]: """Feed a text chunk, return (thinking_delta, content_delta). Args: text: New text chunk from model output. Returns: Tuple of (thinking_text, content_text) extracted from this chunk. """ if not text: return ("", "") # Prepend any buffered partial tag content text = self._buffer + text self._buffer = "" thinking_out = [] content_out = [] i = 0 while i < len(text): if text[i] == '<': # Check if this could be a tag start remaining = text[i:] # Try to match if remaining.startswith(_OPEN_TAG): self._in_thinking = True i += _OPEN_LEN continue if remaining.startswith(_HY3_OPEN_TAG): self._in_thinking = True i += len(_HY3_OPEN_TAG) continue # Try to match if remaining.startswith(_CLOSE_TAG): self._in_thinking = False self._close_seen = True i += _CLOSE_LEN continue if remaining.startswith(_HY3_CLOSE_TAG): self._in_thinking = False self._close_seen = True i += len(_HY3_CLOSE_TAG) continue # Check if it could be a partial tag (not enough chars yet) if self._could_be_tag(remaining): # Buffer the rest and wait for more data self._buffer = remaining break # Not a tag, emit the '<' as regular content if self._in_thinking: thinking_out.append('<') else: content_out.append('<') i += 1 else: if self._in_thinking: thinking_out.append(text[i]) else: content_out.append(text[i]) i += 1 thinking_delta = "".join(thinking_out) content_delta = "".join(content_out) if thinking_delta: self._thinking_accumulated.append(thinking_delta) if content_delta: self._content_emitted = True return (thinking_delta, content_delta) def finish(self) -> Tuple[str, str]: """Flush any remaining buffered content. Should be called when the stream is complete to emit any buffered characters that were waiting for potential tag completion. Also recovers from malformed thinking — when the model never emitted ```` and no content was ever produced, returns the accumulated thinking text as content so the client surfaces a non-empty answer body. Returns: Tuple of (thinking_text, content_text) from remaining buffer (plus recovered content if applicable). """ partial = self._buffer self._buffer = "" # Recovery: prompt opened a thinking block (or model echoed # ```` itself), the close tag never arrived, and nothing # ever streamed as content. Re-emit the accumulated thinking text # as content so the answer body is not empty. The thinking events # already streamed live cannot be retracted, so the client sees # the same text twice — once in the thinking panel, once as the # answer. UX trade-off documented in the chat template plan. if ( self._in_thinking and not self._close_seen and not self._content_emitted and self._thinking_accumulated ): recovered = "".join(self._thinking_accumulated) + partial self._content_emitted = True return ("", recovered) if not partial: return ("", "") # Partial tag never completed — emit it as-is in the current mode. if self._in_thinking: self._thinking_accumulated.append(partial) return (partial, "") else: self._content_emitted = True return ("", partial) @staticmethod def _could_be_tag(text: str) -> bool: """Check if text could be the start of a or tag. Returns True if text is a proper prefix of either tag but not yet a complete match. """ length = len(text) if length >= len(_HY3_CLOSE_TAG): # Long enough to determine - not a partial tag return False # Check against all recognised tags for tag in (_OPEN_TAG, _CLOSE_TAG, _HY3_OPEN_TAG, _HY3_CLOSE_TAG): if length < len(tag) and tag[:length] == text: return True return False class ThinkingBudgetProcessor: """Logits processor that enforces a thinking token budget. Counts tokens generated while in thinking mode. When the budget is exceeded, forces the close-think token(s) one at a time, then becomes a no-op for the rest of generation. Handles both single-token and multi-token close-think sequences, and supports alternative think markers (e.g. ````). Args: think_end_token_ids: Token ID(s) for the close-think tag. budget: Maximum number of thinking tokens before forcing close. think_start_token_id: Token ID for the open-think tag (re-entry detection). """ def __init__( self, think_end_token_ids: List[int], budget: int, think_start_token_id: Optional[int] = None, leading_token_ids: Optional[List[int]] = None, trailing_token_ids: Optional[List[int]] = None, token_to_piece: Optional[Callable[[int], str | bytes | None]] = None, ): self._think_end_ids = think_end_token_ids # Full force sequence: \n + + \n\n (matches training pattern) self._force_sequence = ( (leading_token_ids or []) + list(think_end_token_ids) + (trailing_token_ids or []) ) self._budget = budget self._think_start_id = think_start_token_id self._token_to_piece = token_to_piece # State self._thinking_tokens: int = 0 self._in_thinking: bool = True # Starts True (prompt ends with ) self._forcing: bool = False self._waiting_utf8: bool = False self._force_idx: int = 0 self._done: bool = False self._first_call: bool = True # Sliding window for multi-token end detection self._recent_tokens: List[int] = [] self._last_token_utf8_complete: bool = True self._pending_utf8: bytes = b"" def __call__(self, tokens, logits): """mlx-lm logits processor: (tokens, logits) -> logits.""" import mlx.core as mx # In new mlx-lm API, tokens is the full history list. # Accept each genuinely generated token exactly once (see grammar.py). n = len(tokens) if not hasattr(self, "_accepted_up_to"): self._accepted_up_to = n # skip prompt tokens elif n > self._accepted_up_to: for i in range(self._accepted_up_to, n): self._update_state(int(tokens[i])) self._accepted_up_to = n # If state changed by _update_state, handle immediately if self._done: return logits if self._forcing: return self._force_next_token(logits, mx) if self._waiting_utf8: return logits if self._in_thinking: self._thinking_tokens += 1 if self._thinking_tokens >= self._budget: if self._last_token_utf8_complete: self._forcing = True self._force_idx = 0 self._recent_tokens = [] return self._force_next_token(logits, mx) self._waiting_utf8 = True self._recent_tokens = [] return logits def _update_state(self, token_id: int) -> None: """Update thinking state based on the last generated token.""" self._last_token_utf8_complete = self._is_utf8_complete(token_id) if self._done: if self._think_start_id and token_id == self._think_start_id: self._in_thinking = True self._done = False self._thinking_tokens = 0 self._recent_tokens = [] return if self._forcing: self._force_idx += 1 if self._force_idx >= len(self._force_sequence): self._in_thinking = False self._forcing = False self._done = True return # Detect natural close-think via sliding window if len(self._think_end_ids) == 1: if token_id == self._think_end_ids[0]: self._in_thinking = False self._done = True return else: self._recent_tokens.append(token_id) if len(self._recent_tokens) > len(self._think_end_ids): self._recent_tokens.pop(0) if self._recent_tokens == self._think_end_ids: self._in_thinking = False self._done = True return if self._waiting_utf8: if self._last_token_utf8_complete: self._waiting_utf8 = False self._forcing = True self._force_idx = 0 self._recent_tokens = [] return # Detect re-entry into thinking (rare but possible) if not self._in_thinking and self._think_start_id and token_id == self._think_start_id: self._in_thinking = True self._done = False self._thinking_tokens = 0 self._recent_tokens = [] def _is_utf8_complete(self, token_id: int) -> bool: """Best-effort UTF-8 boundary check for accepted token bytes.""" if self._token_to_piece is None: return True try: piece = self._token_to_piece(token_id) except Exception: return True if piece is None: return True if isinstance(piece, str): self._pending_utf8 = b"" return True self._pending_utf8 += piece try: self._pending_utf8.decode("utf-8") self._pending_utf8 = b"" return True except UnicodeDecodeError as exc: if exc.reason == "unexpected end of data" or exc.end == len( self._pending_utf8 ): return False self._pending_utf8 = b"" return True def _force_next_token(self, logits, mx): """Force the next token in the close-think + trailing sequence.""" target_id = self._force_sequence[self._force_idx] forced = mx.full(logits.shape, float("-inf")) forced[..., target_id] = 0.0 return forced