Add GLM45 tool interruption support (#17714)
Co-authored-by: Sisyphus <clio-agent@sisyphuslabs.ai>
This commit is contained in:
@@ -25,11 +25,13 @@ class BaseReasoningFormatDetector:
|
|||||||
think_end_token: str,
|
think_end_token: str,
|
||||||
force_reasoning: bool = False,
|
force_reasoning: bool = False,
|
||||||
stream_reasoning: bool = True,
|
stream_reasoning: bool = True,
|
||||||
|
tool_start_token: Optional[str] = None,
|
||||||
continue_final_message: bool = False,
|
continue_final_message: bool = False,
|
||||||
previous_content: str = "",
|
previous_content: str = "",
|
||||||
):
|
):
|
||||||
self.think_start_token = think_start_token
|
self.think_start_token = think_start_token
|
||||||
self.think_end_token = think_end_token
|
self.think_end_token = think_end_token
|
||||||
|
self.tool_start_token = tool_start_token
|
||||||
self._in_reasoning = force_reasoning
|
self._in_reasoning = force_reasoning
|
||||||
self.stream_reasoning = stream_reasoning
|
self.stream_reasoning = stream_reasoning
|
||||||
|
|
||||||
@@ -66,7 +68,21 @@ class BaseReasoningFormatDetector:
|
|||||||
self.think_end_token not in processed_text
|
self.think_end_token not in processed_text
|
||||||
and self.think_end_token not in self.previous_content
|
and self.think_end_token not in self.previous_content
|
||||||
):
|
):
|
||||||
# Assume reasoning was truncated before `</think>` token
|
# Check for tool_start_token interruption
|
||||||
|
if (
|
||||||
|
in_reasoning
|
||||||
|
and self.tool_start_token is not None
|
||||||
|
and self.tool_start_token in processed_text
|
||||||
|
):
|
||||||
|
# Find the first occurrence of tool_start_token and split there
|
||||||
|
tool_idx = processed_text.find(self.tool_start_token)
|
||||||
|
reasoning_text = processed_text[:tool_idx].strip()
|
||||||
|
# Preserve tool_start_token in normal text
|
||||||
|
normal_text = processed_text[tool_idx:]
|
||||||
|
return StreamingParseResult(
|
||||||
|
normal_text=normal_text, reasoning_text=reasoning_text
|
||||||
|
)
|
||||||
|
# Assume reasoning was truncated before end token
|
||||||
return StreamingParseResult(reasoning_text=processed_text)
|
return StreamingParseResult(reasoning_text=processed_text)
|
||||||
|
|
||||||
# Extract reasoning content
|
# Extract reasoning content
|
||||||
@@ -96,9 +112,12 @@ class BaseReasoningFormatDetector:
|
|||||||
current_text = self._buffer
|
current_text = self._buffer
|
||||||
|
|
||||||
# If the current text is a prefix of the think token, keep buffering
|
# If the current text is a prefix of the think token, keep buffering
|
||||||
|
tokens_to_check = [self.think_start_token, self.think_end_token]
|
||||||
|
if self.tool_start_token:
|
||||||
|
tokens_to_check.append(self.tool_start_token)
|
||||||
if any(
|
if any(
|
||||||
token.startswith(current_text) and token != current_text
|
token.startswith(current_text) and token != current_text
|
||||||
for token in [self.think_start_token, self.think_end_token]
|
for token in tokens_to_check
|
||||||
):
|
):
|
||||||
return StreamingParseResult()
|
return StreamingParseResult()
|
||||||
|
|
||||||
@@ -124,6 +143,17 @@ class BaseReasoningFormatDetector:
|
|||||||
|
|
||||||
# Continue with reasoning content
|
# Continue with reasoning content
|
||||||
if self._in_reasoning:
|
if self._in_reasoning:
|
||||||
|
# Check for tool_start_token interruption
|
||||||
|
if self.tool_start_token and self.tool_start_token in current_text:
|
||||||
|
tool_idx = current_text.find(self.tool_start_token)
|
||||||
|
reasoning_text = current_text[:tool_idx]
|
||||||
|
# Preserve tool_start_token in normal text
|
||||||
|
normal_text = current_text[tool_idx:]
|
||||||
|
self._buffer = ""
|
||||||
|
self._in_reasoning = False
|
||||||
|
return StreamingParseResult(
|
||||||
|
normal_text=normal_text, reasoning_text=reasoning_text
|
||||||
|
)
|
||||||
if self.stream_reasoning:
|
if self.stream_reasoning:
|
||||||
# Stream the content immediately
|
# Stream the content immediately
|
||||||
self._buffer = ""
|
self._buffer = ""
|
||||||
@@ -238,6 +268,29 @@ class KimiDetector(BaseReasoningFormatDetector):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class Glm45Detector(BaseReasoningFormatDetector):
|
||||||
|
"""
|
||||||
|
Detector for GLM-4.5 models.
|
||||||
|
Assumes reasoning format:
|
||||||
|
(<think>)*(.*)</think>
|
||||||
|
|
||||||
|
GLM-4.5 uses `<tool_call>` as the tool start token to switch from reasoning mode to normal mode.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
stream_reasoning (bool): If False, accumulates reasoning content until the end tag.
|
||||||
|
If True, streams reasoning content as it arrives.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, stream_reasoning: bool = True, force_reasoning: bool = False):
|
||||||
|
super().__init__(
|
||||||
|
"<think>",
|
||||||
|
"</think>",
|
||||||
|
force_reasoning=force_reasoning,
|
||||||
|
stream_reasoning=stream_reasoning,
|
||||||
|
tool_start_token="<tool_call>",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class GptOssDetector(BaseReasoningFormatDetector):
|
class GptOssDetector(BaseReasoningFormatDetector):
|
||||||
"""
|
"""
|
||||||
Detector for T4-style reasoning format (GPT-OSS), using the HarmonyParser.
|
Detector for T4-style reasoning format (GPT-OSS), using the HarmonyParser.
|
||||||
@@ -375,7 +428,7 @@ class ReasoningParser:
|
|||||||
DetectorMap: Dict[str, Type[BaseReasoningFormatDetector]] = {
|
DetectorMap: Dict[str, Type[BaseReasoningFormatDetector]] = {
|
||||||
"deepseek-r1": DeepSeekR1Detector,
|
"deepseek-r1": DeepSeekR1Detector,
|
||||||
"deepseek-v3": Qwen3Detector,
|
"deepseek-v3": Qwen3Detector,
|
||||||
"glm45": Qwen3Detector,
|
"glm45": Glm45Detector,
|
||||||
"gpt-oss": GptOssDetector,
|
"gpt-oss": GptOssDetector,
|
||||||
"kimi": KimiDetector,
|
"kimi": KimiDetector,
|
||||||
"kimi_k2": Qwen3Detector,
|
"kimi_k2": Qwen3Detector,
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ import unittest
|
|||||||
from sglang.srt.parser.reasoning_parser import (
|
from sglang.srt.parser.reasoning_parser import (
|
||||||
BaseReasoningFormatDetector,
|
BaseReasoningFormatDetector,
|
||||||
DeepSeekR1Detector,
|
DeepSeekR1Detector,
|
||||||
|
Glm45Detector,
|
||||||
KimiDetector,
|
KimiDetector,
|
||||||
Qwen3Detector,
|
Qwen3Detector,
|
||||||
ReasoningParser,
|
ReasoningParser,
|
||||||
@@ -313,6 +314,158 @@ class TestKimiDetector(CustomTestCase):
|
|||||||
self.assertEqual(result.normal_text, "answer")
|
self.assertEqual(result.normal_text, "answer")
|
||||||
|
|
||||||
|
|
||||||
|
class TestGlm45Detector(CustomTestCase):
|
||||||
|
"""Test cases for GLM45 detector with tool interruption support."""
|
||||||
|
|
||||||
|
def setUp(self):
|
||||||
|
self.detector = Glm45Detector()
|
||||||
|
|
||||||
|
def test_init(self):
|
||||||
|
"""Test Glm45Detector initialization."""
|
||||||
|
self.assertEqual(self.detector.think_start_token, "<think>")
|
||||||
|
self.assertEqual(self.detector.think_end_token, "</think>")
|
||||||
|
self.assertEqual(self.detector.tool_start_token, "<tool_call>")
|
||||||
|
self.assertFalse(self.detector._in_reasoning)
|
||||||
|
self.assertTrue(self.detector.stream_reasoning)
|
||||||
|
|
||||||
|
def test_detect_and_parse_normal_reasoning(self):
|
||||||
|
"""Test parsing normal reasoning block without tool interruption."""
|
||||||
|
text = "<think>Let me think about this step by step</think>The answer is 42."
|
||||||
|
result = self.detector.detect_and_parse(text)
|
||||||
|
self.assertEqual(result.reasoning_text, "Let me think about this step by step")
|
||||||
|
self.assertEqual(result.normal_text, "The answer is 42.")
|
||||||
|
|
||||||
|
def test_detect_and_parse_tool_interrupt(self):
|
||||||
|
"""
|
||||||
|
Test parsing with tool interruption.
|
||||||
|
|
||||||
|
GLM45 can interrupt reasoning with tool token (<tool_call>) without closing </think>.
|
||||||
|
Should split at the first occurrence of tool_start_token using find().
|
||||||
|
"""
|
||||||
|
text = "<think>I need to think<tool_call>tool call data"
|
||||||
|
result = self.detector.detect_and_parse(text)
|
||||||
|
self.assertEqual(result.reasoning_text, "I need to think")
|
||||||
|
self.assertEqual(result.normal_text, "<tool_call>tool call data")
|
||||||
|
|
||||||
|
def test_detect_and_parse_multiple_tool_calls_find(self):
|
||||||
|
"""
|
||||||
|
Test that find() finds the FIRST occurrence of tool_start_token.
|
||||||
|
|
||||||
|
If multiple tool calls exist in buffer, should split at the first one.
|
||||||
|
"""
|
||||||
|
text = "<think>thinking<tool_call>first tool<tool_call>second tool<tool_call>final tool"
|
||||||
|
result = self.detector.detect_and_parse(text)
|
||||||
|
# Should split at the first <tool_call>
|
||||||
|
self.assertEqual(result.reasoning_text, "thinking")
|
||||||
|
self.assertEqual(
|
||||||
|
result.normal_text,
|
||||||
|
"<tool_call>first tool<tool_call>second tool<tool_call>final tool",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_detect_and_parse_truncated_reasoning(self):
|
||||||
|
"""
|
||||||
|
Test truncated reasoning without tool or end tag.
|
||||||
|
|
||||||
|
Should return all content as reasoning_text.
|
||||||
|
"""
|
||||||
|
text = "<think>This is incomplete"
|
||||||
|
result = self.detector.detect_and_parse(text)
|
||||||
|
self.assertEqual(result.reasoning_text, "This is incomplete")
|
||||||
|
self.assertEqual(result.normal_text, "")
|
||||||
|
|
||||||
|
def test_detect_and_parse_normal_text_only(self):
|
||||||
|
"""Test parsing text without reasoning block."""
|
||||||
|
text = "Just the answer without any reasoning."
|
||||||
|
result = self.detector.detect_and_parse(text)
|
||||||
|
self.assertEqual(result.normal_text, text)
|
||||||
|
self.assertEqual(result.reasoning_text, "")
|
||||||
|
|
||||||
|
def test_streaming_normal_flow(self):
|
||||||
|
"""Test streaming with normal reasoning flow."""
|
||||||
|
# Start reasoning
|
||||||
|
result1 = self.detector.parse_streaming_increment("<think>")
|
||||||
|
self.assertEqual(result1.normal_text, "")
|
||||||
|
self.assertEqual(result1.reasoning_text, "")
|
||||||
|
self.assertTrue(self.detector._in_reasoning)
|
||||||
|
|
||||||
|
# Reasoning content
|
||||||
|
result2 = self.detector.parse_streaming_increment("thinking...")
|
||||||
|
self.assertEqual(result2.normal_text, "")
|
||||||
|
self.assertEqual(result2.reasoning_text, "thinking...")
|
||||||
|
|
||||||
|
# End reasoning
|
||||||
|
result3 = self.detector.parse_streaming_increment("</think>answer")
|
||||||
|
self.assertEqual(result3.normal_text, "answer")
|
||||||
|
self.assertEqual(result3.reasoning_text, "")
|
||||||
|
self.assertFalse(self.detector._in_reasoning)
|
||||||
|
|
||||||
|
def test_streaming_tool_interrupt_split_tokens(self):
|
||||||
|
"""
|
||||||
|
Test streaming with tool interruption where tool token is split across chunks.
|
||||||
|
|
||||||
|
This tests the buffer prefix logic that prevents partial emission of tool token.
|
||||||
|
"""
|
||||||
|
# Start reasoning
|
||||||
|
self.detector.parse_streaming_increment("<think>")
|
||||||
|
|
||||||
|
# Add reasoning
|
||||||
|
result1 = self.detector.parse_streaming_increment("thinking")
|
||||||
|
self.assertEqual(result1.reasoning_text, "thinking")
|
||||||
|
|
||||||
|
# Send partial tool token (should be buffered, not emitted)
|
||||||
|
result2 = self.detector.parse_streaming_increment("<tool_call>")
|
||||||
|
# Tool token is in buffer, causing switch to normal mode
|
||||||
|
self.assertEqual(result2.reasoning_text, "")
|
||||||
|
self.assertEqual(result2.normal_text, "<tool_call>")
|
||||||
|
self.assertFalse(self.detector._in_reasoning)
|
||||||
|
|
||||||
|
# Send tool args
|
||||||
|
result3 = self.detector.parse_streaming_increment("tool args")
|
||||||
|
self.assertEqual(result3.reasoning_text, "")
|
||||||
|
self.assertEqual(result3.normal_text, "tool args")
|
||||||
|
|
||||||
|
def test_streaming_no_stream_reasoning(self):
|
||||||
|
"""Test streaming without stream_reasoning enabled."""
|
||||||
|
detector = Glm45Detector(stream_reasoning=False)
|
||||||
|
|
||||||
|
# Start reasoning
|
||||||
|
detector.parse_streaming_increment("<think>")
|
||||||
|
|
||||||
|
# Reasoning content is buffered and not returned yet
|
||||||
|
result = detector.parse_streaming_increment("thinking")
|
||||||
|
self.assertEqual(result.reasoning_text, "")
|
||||||
|
self.assertEqual(result.normal_text, "")
|
||||||
|
|
||||||
|
# Tool interruption should still work - flushes buffered reasoning
|
||||||
|
# Note: buffer preserves original text including <think> tag
|
||||||
|
result = detector.parse_streaming_increment("<tool_call>tool call")
|
||||||
|
self.assertEqual(result.reasoning_text, "<think>thinking")
|
||||||
|
self.assertEqual(result.normal_text, "<tool_call>tool call")
|
||||||
|
|
||||||
|
def test_streaming_empty_reasoning_with_tool(self):
|
||||||
|
"""Test empty reasoning block followed by tool call."""
|
||||||
|
result1 = self.detector.parse_streaming_increment("<think>")
|
||||||
|
result2 = self.detector.parse_streaming_increment("<tool_call>tool call")
|
||||||
|
self.assertEqual(result2.reasoning_text, "")
|
||||||
|
self.assertEqual(result2.normal_text, "<tool_call>tool call")
|
||||||
|
|
||||||
|
def test_forced_reasoning_mode(self):
|
||||||
|
"""Test GLM45 with force_reasoning=True."""
|
||||||
|
detector = Glm45Detector(force_reasoning=True)
|
||||||
|
|
||||||
|
# Without start token, should still be in reasoning mode
|
||||||
|
text = "This is reasoning"
|
||||||
|
result = detector.detect_and_parse(text)
|
||||||
|
self.assertEqual(result.reasoning_text, "This is reasoning")
|
||||||
|
self.assertEqual(result.normal_text, "")
|
||||||
|
|
||||||
|
# Tool interruption should work with forced reasoning
|
||||||
|
text = "More reasoning<tool_call>tool call"
|
||||||
|
result = detector.detect_and_parse(text)
|
||||||
|
self.assertEqual(result.reasoning_text, "More reasoning")
|
||||||
|
self.assertEqual(result.normal_text, "<tool_call>tool call")
|
||||||
|
|
||||||
|
|
||||||
class TestReasoningParser(CustomTestCase):
|
class TestReasoningParser(CustomTestCase):
|
||||||
def test_init_valid_model(self):
|
def test_init_valid_model(self):
|
||||||
"""Test initialization with valid model types."""
|
"""Test initialization with valid model types."""
|
||||||
@@ -325,6 +478,9 @@ class TestReasoningParser(CustomTestCase):
|
|||||||
parser = ReasoningParser("kimi")
|
parser = ReasoningParser("kimi")
|
||||||
self.assertIsInstance(parser.detector, KimiDetector)
|
self.assertIsInstance(parser.detector, KimiDetector)
|
||||||
|
|
||||||
|
parser = ReasoningParser("glm45")
|
||||||
|
self.assertIsInstance(parser.detector, Glm45Detector)
|
||||||
|
|
||||||
def test_init_invalid_model(self):
|
def test_init_invalid_model(self):
|
||||||
"""Test initialization with invalid model type."""
|
"""Test initialization with invalid model type."""
|
||||||
with self.assertRaises(ValueError) as context:
|
with self.assertRaises(ValueError) as context:
|
||||||
@@ -383,6 +539,32 @@ class TestReasoningParser(CustomTestCase):
|
|||||||
parser = ReasoningParser("qwen3", stream_reasoning=True)
|
parser = ReasoningParser("qwen3", stream_reasoning=True)
|
||||||
self.assertTrue(parser.detector.stream_reasoning)
|
self.assertTrue(parser.detector.stream_reasoning)
|
||||||
|
|
||||||
|
def test_glm45_tool_interruption(self):
|
||||||
|
"""Test GLM45 tool interruption through ReasoningParser API."""
|
||||||
|
parser = ReasoningParser("glm45")
|
||||||
|
|
||||||
|
# Non-streaming: tool interrupt
|
||||||
|
reasoning, normal = parser.parse_non_stream(
|
||||||
|
"<think>thinking<tool_call>tool call"
|
||||||
|
)
|
||||||
|
self.assertEqual(reasoning, "thinking")
|
||||||
|
self.assertEqual(normal, "<tool_call>tool call")
|
||||||
|
|
||||||
|
# Streaming: tool interrupt
|
||||||
|
parser = ReasoningParser("glm45")
|
||||||
|
chunks = ["<think>", "reasoning", "<tool_call>", "tool args"]
|
||||||
|
all_reasoning = ""
|
||||||
|
all_normal = ""
|
||||||
|
for chunk in chunks:
|
||||||
|
reasoning, normal = parser.parse_stream_chunk(chunk)
|
||||||
|
if reasoning:
|
||||||
|
all_reasoning += reasoning
|
||||||
|
if normal:
|
||||||
|
all_normal += normal
|
||||||
|
|
||||||
|
self.assertEqual(all_reasoning, "reasoning")
|
||||||
|
self.assertEqual(all_normal, "<tool_call>tool args")
|
||||||
|
|
||||||
|
|
||||||
class TestIntegrationScenarios(CustomTestCase):
|
class TestIntegrationScenarios(CustomTestCase):
|
||||||
"""Integration tests for realistic usage scenarios."""
|
"""Integration tests for realistic usage scenarios."""
|
||||||
|
|||||||
Reference in New Issue
Block a user