Migrate parsing tests to test/registered/parser/ (#16467)

This commit is contained in:
Alison Shao
2026-01-05 21:15:32 -08:00
committed by GitHub
parent 3aa11ca722
commit 17958c5f0d
6 changed files with 52 additions and 5 deletions
-4
View File
@@ -43,7 +43,6 @@ suites = {
TestFile("test_fp8_utils.py", 9),
TestFile("rotary_embedding/test_mrope.py", 10),
TestFile("test_gpt_oss_1gpu.py", 402),
TestFile("test_harmony_parser.py", 6),
TestFile("test_hidden_states.py", 55),
TestFile("test_input_embeddings.py", 38),
TestFile("test_io_struct.py", 8),
@@ -62,7 +61,6 @@ suites = {
TestFile("test_prefill_adder.py", 1),
TestFile("test_priority_scheduling.py", 130),
TestFile("test_pytorch_sampling_backend.py", 66),
TestFile("test_reasoning_parser.py", 5),
TestFile("test_request_queue_validation.py", 47),
TestFile("test_retract_decode.py", 259),
TestFile("test_score_api.py", 260),
@@ -207,7 +205,6 @@ suite_amd = {
TestFile("test_constrained_decoding.py", 120),
TestFile("test_eval_fp8_accuracy.py", 303),
TestFile("test_external_models.py", 45),
TestFile("test_harmony_parser.py", 20),
TestFile("test_input_embeddings.py", 38),
TestFile("test_io_struct.py", 8),
TestFile("test_jinja_template_utils.py", 1),
@@ -224,7 +221,6 @@ suite_amd = {
TestFile("test_profile_merger.py", 12),
TestFile("test_profile_merger_http_api.py", 15),
TestFile("test_pytorch_sampling_backend.py", 66),
TestFile("test_reasoning_parser.py", 5),
TestFile("test_request_queue_validation.py", 70),
TestFile("test_retract_decode.py", 450),
TestFile("test_rope_rocm.py", 3),
-876
View File
@@ -1,876 +0,0 @@
import unittest
from sglang.srt.parser.harmony_parser import (
CanonicalStrategy,
Event,
HarmonyParser,
TextStrategy,
Token,
iter_tokens,
prefix_hold,
)
from sglang.test.test_utils import CustomTestCase
class TestEvent(CustomTestCase):
def test_init(self):
"""Test Event dataclass initialization."""
event = Event("reasoning", "content")
self.assertEqual(event.event_type, "reasoning")
self.assertEqual(event.content, "content")
class TestToken(CustomTestCase):
def test_init(self):
"""Test Token dataclass initialization."""
token = Token("START", 0, 7)
self.assertEqual(token.type, "START")
self.assertEqual(token.start, 0)
self.assertEqual(token.end, 7)
class TestPrefixHold(CustomTestCase):
def test_empty_text(self):
"""Test prefix_hold with empty text."""
emit, hold = prefix_hold("", ["<|start|>"])
self.assertEqual(emit, "")
self.assertEqual(hold, "")
def test_no_matching_prefixes(self):
"""Test prefix_hold with no matching prefixes."""
emit, hold = prefix_hold("hello world", ["<|start|>", "<|end|>"])
self.assertEqual(emit, "hello world")
self.assertEqual(hold, "")
def test_partial_token_suffix(self):
"""Test prefix_hold with partial token at end."""
emit, hold = prefix_hold("hello <|ret", ["<|return|>"])
self.assertEqual(emit, "hello ")
self.assertEqual(hold, "<|ret")
def test_multiple_potential_matches(self):
"""Test prefix_hold with multiple potential matches."""
emit, hold = prefix_hold("text <|", ["<|start|>", "<|end|>"])
self.assertEqual(emit, "text ")
self.assertEqual(hold, "<|")
def test_exact_token_match(self):
"""Test prefix_hold with exact token match."""
emit, hold = prefix_hold("text <|start|>", ["<|start|>"])
self.assertEqual(emit, "text <|start|>")
self.assertEqual(hold, "")
class TestIterTokens(CustomTestCase):
def test_empty_text(self):
"""Test iter_tokens with empty text."""
tokens = list(iter_tokens(""))
self.assertEqual(tokens, [])
def test_plain_text(self):
"""Test iter_tokens with plain text."""
tokens = list(iter_tokens("hello world"))
self.assertEqual(len(tokens), 1)
self.assertEqual(tokens[0].type, "TEXT")
self.assertEqual(tokens[0].start, 0)
self.assertEqual(tokens[0].end, 11)
def test_single_token(self):
"""Test iter_tokens with single structural token."""
tokens = list(iter_tokens("<|start|>"))
self.assertEqual(len(tokens), 1)
self.assertEqual(tokens[0].type, "START")
self.assertEqual(tokens[0].start, 0)
self.assertEqual(tokens[0].end, 9)
def test_mixed_content(self):
"""Test iter_tokens with mixed text and tokens."""
tokens = list(iter_tokens("text<|start|>more text"))
self.assertEqual(len(tokens), 3)
self.assertEqual(tokens[0].type, "TEXT")
self.assertEqual(tokens[0].start, 0)
self.assertEqual(tokens[0].end, 4)
self.assertEqual(tokens[1].type, "START")
self.assertEqual(tokens[1].start, 4)
self.assertEqual(tokens[1].end, 13)
self.assertEqual(tokens[2].type, "TEXT")
self.assertEqual(tokens[2].start, 13)
self.assertEqual(tokens[2].end, 22)
def test_unknown_token_partial_suffix(self):
"""Test iter_tokens with unknown token that could be partial."""
tokens = list(iter_tokens("text <|ret"))
self.assertEqual(len(tokens), 2)
self.assertEqual(tokens[0].type, "TEXT")
self.assertEqual(tokens[0].start, 0)
self.assertEqual(tokens[0].end, 5)
self.assertEqual(tokens[1].type, "TEXT")
self.assertEqual(tokens[1].start, 5)
self.assertEqual(tokens[1].end, 10)
def test_unknown_token_middle(self):
"""Test iter_tokens with unknown token in middle."""
tokens = list(iter_tokens("text <|weird|> more <|start|>"))
self.assertEqual(len(tokens), 5)
self.assertEqual(tokens[0].type, "TEXT")
self.assertEqual(tokens[1].type, "TEXT") # "<|"
self.assertEqual(tokens[2].type, "TEXT") # "weird|> more "
self.assertEqual(tokens[3].type, "START")
# No trailing text token since it ends with a known token
def test_all_structural_tokens(self):
"""Test iter_tokens recognizes all structural tokens."""
text = "<|start|><|channel|><|message|><|constrain|><|end|><|call|><|return|>"
tokens = list(iter_tokens(text))
expected_types = [
"START",
"CHANNEL",
"MESSAGE",
"CONSTRAIN",
"END",
"CALL",
"RETURN",
]
self.assertEqual(len(tokens), len(expected_types))
for token, expected_type in zip(tokens, expected_types):
self.assertEqual(token.type, expected_type)
class TestCanonicalStrategy(CustomTestCase):
def setUp(self):
self.strategy = CanonicalStrategy()
def test_init(self):
"""Test CanonicalStrategy initialization."""
self.assertIn("<|start|>", self.strategy.guard_tokens)
self.assertIn("<|constrain|>", self.strategy.guard_tokens)
def test_extract_channel_type(self):
"""Test _extract_channel_type method."""
self.assertEqual(self.strategy._extract_channel_type("analysis"), "analysis")
self.assertEqual(
self.strategy._extract_channel_type("commentary to=functions.tool"),
"commentary",
)
self.assertEqual(self.strategy._extract_channel_type("final to=user"), "final")
self.assertEqual(self.strategy._extract_channel_type("ANALYSIS"), "analysis")
self.assertIsNone(self.strategy._extract_channel_type("unknown"))
def test_parse_single_analysis_block(self):
"""Test parsing single analysis block."""
text = "<|channel|>analysis<|message|>Let me think about this<|end|>"
events, remaining = self.strategy.parse(text)
self.assertEqual(len(events), 1)
self.assertEqual(events[0].event_type, "reasoning")
self.assertEqual(events[0].content, "Let me think about this")
self.assertEqual(remaining, "")
def test_parse_single_commentary_block(self):
"""Test parsing single commentary block."""
text = "<|channel|>commentary<|message|>User-visible message<|end|>"
events, remaining = self.strategy.parse(text)
self.assertEqual(len(events), 1)
self.assertEqual(events[0].event_type, "normal")
self.assertEqual(events[0].content, "User-visible message")
self.assertEqual(remaining, "")
def test_parse_single_final_block(self):
"""Test parsing single final block."""
text = "<|start|>assistant<|channel|>final<|message|>The answer is 42<|return|>"
events, remaining = self.strategy.parse(text)
self.assertEqual(len(events), 1)
self.assertEqual(events[0].event_type, "normal")
self.assertEqual(events[0].content, "The answer is 42")
self.assertEqual(remaining, "")
def test_parse_tool_call_commentary(self):
"""Test parsing tool call on commentary channel."""
text = '<|channel|>commentary to=functions.get_weather<|message|>{"location": "SF"}<|call|>'
events, remaining = self.strategy.parse(text)
self.assertEqual(len(events), 1)
self.assertEqual(events[0].event_type, "tool_call")
self.assertEqual(events[0].content, '{"location": "SF"}')
self.assertEqual(remaining, "")
def test_parse_tool_call_analysis(self):
"""Test parsing built-in tool call on analysis channel."""
text = '<|channel|>analysis to=browser.search<|message|>{"query": "SGLang"}<|call|>'
events, remaining = self.strategy.parse(text)
self.assertEqual(len(events), 1)
self.assertEqual(events[0].event_type, "tool_call")
self.assertEqual(events[0].content, '{"query": "SGLang"}')
self.assertEqual(remaining, "")
def test_parse_complex_sequence(self):
"""Test parsing complex sequence with multiple blocks."""
text = (
"<|channel|>analysis<|message|>Need to use function get_weather.<|end|>"
"<|start|>assistant<|channel|>commentary to=functions.get_weather<|message|>"
'{"location":"San Francisco"}<|call|>'
)
events, remaining = self.strategy.parse(text)
self.assertEqual(len(events), 2)
self.assertEqual(events[0].event_type, "reasoning")
self.assertEqual(events[0].content, "Need to use function get_weather.")
self.assertEqual(events[1].event_type, "tool_call")
self.assertEqual(events[1].content, '{"location":"San Francisco"}')
self.assertEqual(remaining, "")
def test_parse_with_interspersed_text(self):
"""Test parsing with plain text between blocks."""
text = (
"Some text "
"<|channel|>analysis<|message|>reasoning<|end|>"
" more text "
"<|start|>assistant<|channel|>final<|message|>answer<|return|>"
" trailing text"
)
events, remaining = self.strategy.parse(text)
self.assertEqual(len(events), 4)
self.assertEqual(events[0].event_type, "normal")
self.assertEqual(events[0].content, "Some text ")
self.assertEqual(events[1].event_type, "reasoning")
self.assertEqual(events[1].content, "reasoning")
self.assertEqual(events[2].event_type, "normal")
self.assertEqual(events[2].content, " more text ")
self.assertEqual(events[3].event_type, "normal")
self.assertEqual(events[3].content, "answer trailing text")
self.assertEqual(remaining, "")
def test_parse_incomplete_block(self):
"""Test parsing incomplete block (streaming scenario)."""
text = "<|channel|>analysis<|message|>partial content"
events, remaining = self.strategy.parse(text)
self.assertEqual(len(events), 1)
self.assertEqual(events[0].event_type, "reasoning")
self.assertEqual(events[0].content, "partial content")
self.assertEqual(remaining, "<|channel|>analysis<|message|>")
def test_parse_partial_token_suffix(self):
"""Test parsing with partial token at end."""
text = "complete text <|ret"
events, remaining = self.strategy.parse(text)
self.assertEqual(len(events), 1)
self.assertEqual(events[0].event_type, "normal")
self.assertEqual(events[0].content, "complete text ")
self.assertEqual(remaining, "<|ret")
def test_parse_tool_response_message(self):
"""Test parsing tool response message (no channel)."""
text = '<|start|>functions.get_weather to=assistant<|message|>{"sunny": true}<|end|>'
events, remaining = self.strategy.parse(text)
self.assertEqual(len(events), 1)
self.assertEqual(events[0].event_type, "normal")
self.assertEqual(events[0].content, '{"sunny": true}')
self.assertEqual(remaining, "")
def test_parse_empty_content_blocks(self):
"""Test parsing blocks with empty content."""
text = "<|channel|>analysis<|message|><|end|>"
events, remaining = self.strategy.parse(text)
self.assertEqual(len(events), 1)
self.assertEqual(events[0].event_type, "reasoning")
self.assertEqual(events[0].content, "")
self.assertEqual(remaining, "")
def test_parse_commentary_filler_between_blocks(self):
"""Test that 'commentary' filler between <|call|> and <|channel|> is filtered out."""
# This pattern occurs when the model generates malformed output
text = (
'<|channel|>commentary to=functions.get_weather<|message|>{"location":"SF"}<|call|>'
"commentary" # This should be filtered out
'<|channel|>commentary to=functions.get_temp<|message|>{"location":"NYC"}<|call|>'
)
events, remaining = self.strategy.parse(text)
# Should have 2 tool calls, no "commentary" normal text
self.assertEqual(len(events), 2)
self.assertEqual(events[0].event_type, "tool_call")
self.assertEqual(events[0].content, '{"location":"SF"}')
self.assertEqual(events[1].event_type, "tool_call")
self.assertEqual(events[1].content, '{"location":"NYC"}')
self.assertEqual(remaining, "")
# Verify no "commentary" text was emitted as normal content
normal_events = [e for e in events if e.event_type == "normal"]
commentary_events = [
e for e in normal_events if "commentary" in e.content.lower()
]
self.assertEqual(
len(commentary_events), 0, "Commentary filler should be filtered out"
)
class TestTextStrategy(CustomTestCase):
def setUp(self):
self.strategy = TextStrategy()
def test_init(self):
"""Test TextStrategy initialization."""
self.assertIn("analysis_then_final", self.strategy.patterns)
def test_parse_analysis_then_final(self):
"""Test parsing analysis then final format."""
text = "analysis I need to think about this. assistantfinal The answer is 42."
events, remaining = self.strategy.parse(text)
self.assertEqual(len(events), 2)
self.assertEqual(events[0].event_type, "reasoning")
self.assertEqual(events[0].content, "I need to think about this.")
self.assertEqual(events[1].event_type, "normal")
self.assertEqual(events[1].content, "The answer is 42.")
self.assertEqual(remaining, "")
def test_parse_commentary_then_final(self):
"""Test parsing commentary then final format."""
text = "commentary User-visible preamble. assistantfinal The answer is 42."
events, remaining = self.strategy.parse(text)
self.assertEqual(len(events), 2)
self.assertEqual(events[0].event_type, "normal")
self.assertEqual(events[0].content, "User-visible preamble.")
self.assertEqual(events[1].event_type, "normal")
self.assertEqual(events[1].content, "The answer is 42.")
self.assertEqual(remaining, "")
def test_parse_final_only(self):
"""Test parsing final-only format."""
text = "assistantfinal The direct answer."
events, remaining = self.strategy.parse(text)
self.assertEqual(len(events), 1)
self.assertEqual(events[0].event_type, "normal")
self.assertEqual(events[0].content, "The direct answer.")
self.assertEqual(remaining, "")
def test_parse_analysis_only(self):
"""Test parsing analysis-only format."""
text = "analysis This is reasoning content."
events, remaining = self.strategy.parse(text)
# For analysis-only, streaming parse should keep header and emit with leading space
self.assertEqual(len(events), 1)
self.assertEqual(events[0].event_type, "reasoning")
self.assertEqual(events[0].content, " This is reasoning content.")
self.assertEqual(remaining, "analysis")
def test_parse_incomplete_assistantfinal(self):
"""Test parsing with incomplete assistantfinal."""
text = "analysis reasoning content assistantfin"
events, remaining = self.strategy.parse(text)
self.assertEqual(len(events), 0)
self.assertEqual(remaining, text) # Hold entire buffer
def test_parse_partial_analysis_streaming(self):
"""Test streaming partial analysis content."""
text = "analysis partial content"
events, remaining = self.strategy.parse(text)
self.assertEqual(len(events), 1)
self.assertEqual(events[0].event_type, "reasoning")
self.assertEqual(events[0].content, " partial content") # Space preserved
self.assertEqual(remaining, "analysis") # Hold header
def test_parse_case_insensitive(self):
"""Test case insensitive parsing."""
text = "ANALYSIS reasoning ASSISTANTFINAL answer"
events, remaining = self.strategy.parse(text)
self.assertEqual(len(events), 2)
self.assertEqual(events[0].event_type, "reasoning")
self.assertEqual(events[1].event_type, "normal")
def test_parse_plain_text_fallback(self):
"""Test parsing plain text without harmony markers."""
text = "Just plain text without any markers."
events, remaining = self.strategy.parse(text)
self.assertEqual(len(events), 1)
self.assertEqual(events[0].event_type, "normal")
self.assertEqual(events[0].content, "Just plain text without any markers.")
self.assertEqual(remaining, "")
def test_parse_analysis_no_space_after_header(self):
"""Test parsing analysis format without space after header (real gpt-oss output)."""
text = "analysisThe user typed random strings. We should respond politely.assistantfinalIt looks like you're testing. How can I help?"
events, remaining = self.strategy.parse(text)
self.assertEqual(len(events), 2)
self.assertEqual(events[0].event_type, "reasoning")
self.assertEqual(
events[0].content,
"The user typed random strings. We should respond politely.",
)
self.assertEqual(events[1].event_type, "normal")
self.assertEqual(
events[1].content, "It looks like you're testing. How can I help?"
)
class TestHarmonyParser(CustomTestCase):
def setUp(self):
self.parser = HarmonyParser()
def test_init(self):
"""Test HarmonyParser initialization."""
self.assertIsNone(self.parser.strategy)
self.assertEqual(self.parser._buffer, "")
def test_strategy_selection_canonical(self):
"""Test automatic strategy selection for canonical format."""
events = self.parser.parse("<|channel|>analysis<|message|>test<|end|>")
self.assertIsInstance(self.parser.strategy, CanonicalStrategy)
self.assertEqual(len(events), 1)
self.assertEqual(events[0].event_type, "reasoning")
def test_strategy_selection_text(self):
"""Test automatic strategy selection for text format."""
events = self.parser.parse("analysis test content")
self.assertIsInstance(self.parser.strategy, TextStrategy)
self.assertEqual(len(events), 1)
self.assertEqual(events[0].event_type, "reasoning")
def test_strategy_selection_delayed(self):
"""Test strategy selection with insufficient initial content."""
# First chunk doesn't have enough info
events1 = self.parser.parse("some")
self.assertEqual(len(events1), 0)
self.assertIsNone(self.parser.strategy)
# Second chunk triggers strategy selection
events2 = self.parser.parse(" analysis content")
self.assertIsInstance(self.parser.strategy, TextStrategy)
self.assertEqual(len(events2), 1)
def test_streaming_canonical_format(self):
"""Test streaming with canonical format."""
chunks = [
"<|channel|>analysis<|message|>",
"reasoning content",
"<|end|>",
"<|start|>assistant<|channel|>final<|message|>",
"final answer",
"<|return|>",
]
all_events = []
for chunk in chunks:
events = self.parser.parse(chunk)
all_events.extend(events)
self.assertEqual(len(all_events), 5)
# Verify we get reasoning events
reasoning_events = [e for e in all_events if e.event_type == "reasoning"]
self.assertTrue(len(reasoning_events) > 0)
# Verify we get normal events
normal_events = [e for e in all_events if e.event_type == "normal"]
self.assertTrue(len(normal_events) > 0)
# Verify content is eventually parsed correctly
combined_reasoning = "".join(e.content for e in reasoning_events)
combined_normal = "".join(
e.content
for e in normal_events
if e.content and "<|return|>" not in e.content
)
self.assertIn("reasoning content", combined_reasoning)
self.assertIn("final answer", combined_normal)
def test_streaming_text_format(self):
"""Test streaming with text format."""
chunks = ["analysis reasoning", " content assistantfinal", " the answer"]
all_events = []
for chunk in chunks:
events = self.parser.parse(chunk)
all_events.extend(events)
# Should have reasoning and normal events
reasoning_events = [e for e in all_events if e.event_type == "reasoning"]
normal_events = [e for e in all_events if e.event_type == "normal"]
self.assertGreater(len(reasoning_events), 0)
self.assertGreater(len(normal_events), 0)
def test_streaming_commentary_filler(self):
"""Test that 'commentary' filler is filtered in streaming case."""
# Test when commentary arrives as a separate chunk after <|call|>
chunks = [
"<|channel|>commentary to=functions.get_weather",
"<|message|>",
'{"location":"SF"}',
"<|call|>",
"comment", # This arrives as separate chunk - should be filtered
"ary", # Continuation of the filler - should be filtered
"<|channel|>commentary to=functions.get_temp",
"<|message|>",
'{"location":"NYC"}',
"<|call|>",
"comment", # Another separate chunk - should be filtered
"ary", # Continuation of the filler - should be filtered
"<|start|>assistant<|channel|>final",
"<|message|>Done<|return|>",
]
all_events = []
for chunk in chunks:
events = self.parser.parse(chunk)
all_events.extend(events)
# Count event types
tool_events = [e for e in all_events if e.event_type == "tool_call"]
normal_events = [e for e in all_events if e.event_type == "normal"]
# Should have 2 tool calls and 1 final message
self.assertEqual(len(tool_events), 2, "Should have 2 tool calls")
self.assertEqual(
len(normal_events), 1, "Should have 1 normal event (final message)"
)
# Verify no "commentary" in normal events
for event in normal_events:
self.assertNotEqual(
event.content.strip().lower(),
"commentary",
"Commentary filler should not appear as normal content in streaming",
)
# Verify content
self.assertEqual(tool_events[0].content, '{"location":"SF"}')
self.assertEqual(tool_events[1].content, '{"location":"NYC"}')
self.assertEqual(normal_events[0].content, "Done")
def test_repetitive_tool_calls_with_commentary_filler(self):
"""Test handling of repetitive tool calls with 'commentary' filler text."""
# This simulates malformed output with repeated tool calls and commentary filler
text = (
"<|channel|>analysis<|message|>Need to get weather<|end|>"
'<|start|>assistant<|channel|>commentary to=functions.get_weather<|message|>{"city":"Boston"}<|call|>'
"commentary" # Filler that should be filtered
'<|channel|>commentary to=functions.get_weather<|message|>{"city":"Boston"}<|call|>'
"commentary" # Another filler
'<|channel|>commentary to=functions.get_weather<|message|>{"city":"Boston"}<|call|>'
"<|channel|>analysis<|message|>Tool not responding<|end|>"
"<|start|>assistant<|channel|>final<|message|>Unable to fetch weather data<|return|>"
)
events = self.parser.parse(text)
# Count event types
reasoning_events = [e for e in events if e.event_type == "reasoning"]
tool_events = [e for e in events if e.event_type == "tool_call"]
normal_events = [e for e in events if e.event_type == "normal"]
# Verify correct number of each type
self.assertEqual(len(reasoning_events), 2, "Should have 2 reasoning events")
self.assertEqual(len(tool_events), 3, "Should have 3 tool calls")
self.assertEqual(
len(normal_events), 1, "Should have 1 normal event (final message)"
)
# Verify no "commentary" filler in normal events
for event in normal_events:
self.assertNotEqual(
event.content.strip().lower(),
"commentary",
"Commentary filler should not appear as normal content",
)
# Verify content is correct
self.assertEqual(reasoning_events[0].content, "Need to get weather")
self.assertEqual(reasoning_events[1].content, "Tool not responding")
self.assertEqual(normal_events[0].content, "Unable to fetch weather data")
class TestIntegrationScenarios(CustomTestCase):
"""Integration tests for realistic Harmony parsing scenarios."""
def test_complete_reasoning_flow(self):
"""Test complete reasoning flow from HARMONY_DOCS.md examples."""
parser = HarmonyParser()
text = (
'<|channel|>analysis<|message|>User asks: "What is 2 + 2?" Simple arithmetic. Provide answer.<|end|>'
"<|start|>assistant<|channel|>final<|message|>2 + 2 = 4.<|return|>"
)
events = parser.parse(text)
self.assertEqual(len(events), 2)
self.assertEqual(events[0].event_type, "reasoning")
self.assertIn("Simple arithmetic", events[0].content)
self.assertEqual(events[1].event_type, "normal")
self.assertEqual(events[1].content, "2 + 2 = 4.")
def test_tool_call_sequence(self):
"""Test tool call sequence from HARMONY_DOCS.md examples."""
parser = HarmonyParser()
text = (
"<|channel|>analysis<|message|>Need to use function get_weather.<|end|>"
"<|start|>assistant<|channel|>commentary to=functions.get_weather <|constrain|>json<|message|>"
'{"location":"San Francisco"}<|call|>'
)
events = parser.parse(text)
self.assertEqual(len(events), 2)
self.assertEqual(events[0].event_type, "reasoning")
self.assertEqual(events[0].content, "Need to use function get_weather.")
self.assertEqual(events[1].event_type, "tool_call")
self.assertEqual(events[1].content, '{"location":"San Francisco"}')
def test_preamble_sequence(self):
"""Test preamble sequence with multiple commentary blocks."""
parser = HarmonyParser()
text = (
"<|channel|>analysis<|message|>Long chain of thought<|end|>"
"<|start|>assistant<|channel|>commentary<|message|>**Action plan**: 1. Generate file 2. Start server<|end|>"
"<|start|>assistant<|channel|>commentary to=functions.generate_file<|message|>"
'{"template": "basic_html"}<|call|>'
)
events = parser.parse(text)
self.assertEqual(len(events), 3)
self.assertEqual(events[0].event_type, "reasoning")
self.assertEqual(events[1].event_type, "normal")
self.assertIn("Action plan", events[1].content)
self.assertEqual(events[2].event_type, "tool_call")
def test_built_in_tool_call(self):
"""Test built-in tool call on analysis channel."""
parser = HarmonyParser()
text = '<|channel|>analysis to=browser.search<|message|>{"query": "SGLang"}<|call|>'
events = parser.parse(text)
self.assertEqual(len(events), 1)
self.assertEqual(events[0].event_type, "tool_call")
self.assertEqual(events[0].content, '{"query": "SGLang"}')
def test_tool_response_handling(self):
"""Test tool response message handling."""
parser = HarmonyParser()
text = '<|start|>functions.get_weather to=assistant<|channel|>commentary<|message|>{"sunny": true, "temperature": 20}<|end|>'
events = parser.parse(text)
self.assertEqual(len(events), 1)
self.assertEqual(events[0].event_type, "normal")
self.assertEqual(events[0].content, '{"sunny": true, "temperature": 20}')
def test_text_fallback_formats(self):
"""Test various text fallback formats."""
parser = HarmonyParser()
# Test analysis then final
events1 = parser.parse("analysis thinking assistantfinal answer")
self.assertEqual(len([e for e in events1 if e.event_type == "reasoning"]), 1)
self.assertEqual(len([e for e in events1 if e.event_type == "normal"]), 1)
# Reset parser for next test
parser = HarmonyParser()
# Test final only
events2 = parser.parse("assistantfinal direct answer")
self.assertEqual(len(events2), 1)
self.assertEqual(events2[0].event_type, "normal")
def test_streaming_property_canonical(self):
"""Test streaming property: chunked parsing produces same semantic content as one-shot parsing."""
full_text = (
"<|channel|>analysis<|message|>reasoning content<|end|>"
"<|start|>assistant<|channel|>final<|message|>final content"
)
# One-shot parsing
parser1 = HarmonyParser()
events_oneshot = parser1.parse(full_text)
events_oneshot += parser1.parse("")
# Chunked parsing
parser2 = HarmonyParser()
chunks = [
"<|channel|>",
"analysis",
"<|message|>",
"reasoning content",
"<|end|>",
"<|start|>assistant",
"<|channel|>final",
"<|message|>",
"final ",
"content",
]
events_chunked = []
for chunk in chunks:
events_chunked.extend(parser2.parse(chunk))
# Compare semantic content rather than exact event structure
reasoning_oneshot = "".join(
e.content for e in events_oneshot if e.event_type == "reasoning"
)
normal_oneshot = "".join(
e.content for e in events_oneshot if e.event_type == "normal"
)
reasoning_chunked = "".join(
e.content for e in events_chunked if e.event_type == "reasoning"
)
normal_chunked = "".join(
e.content for e in events_chunked if e.event_type == "normal"
)
self.assertEqual(reasoning_chunked, reasoning_oneshot)
self.assertEqual(normal_chunked, normal_oneshot)
def test_streaming_property_text(self):
"""Test streaming property for text format."""
full_text = "analysis reasoning content assistantfinal final answer"
# One-shot parsing
parser1 = HarmonyParser()
events_oneshot = parser1.parse(full_text)
# Chunked parsing
parser2 = HarmonyParser()
chunks = ["analysis reason", "ing content assistant", "final final answer"]
events_chunked = []
for chunk in chunks:
events_chunked.extend(parser2.parse(chunk))
# Combine content by type for comparison
reasoning_oneshot = "".join(
e.content for e in events_oneshot if e.event_type == "reasoning"
)
normal_oneshot = "".join(
e.content for e in events_oneshot if e.event_type == "normal"
)
reasoning_chunked = "".join(
e.content for e in events_chunked if e.event_type == "reasoning"
)
normal_chunked = "".join(
e.content for e in events_chunked if e.event_type == "normal"
)
# Account for whitespace differences due to streaming - compare trimmed content
self.assertEqual(reasoning_oneshot.strip(), reasoning_chunked.strip())
self.assertEqual(normal_oneshot.strip(), normal_chunked.strip())
class TestEdgeCases(CustomTestCase):
"""Test edge cases and error conditions."""
def test_malformed_channel_headers(self):
"""Test handling of malformed channel headers."""
parser = HarmonyParser()
# Unknown channel type
text = "<|channel|>unknown<|message|>content<|end|>"
events = parser.parse(text)
# Should be held as incomplete since channel is unknown
self.assertEqual(len(events), 0)
def test_mixed_unknown_tokens(self):
"""Test handling of mixed unknown tokens."""
parser = HarmonyParser()
text = "text <|weird|> more text <|channel|>analysis<|message|>content<|end|>"
events = parser.parse(text)
# Should parse the valid parts
reasoning_events = [e for e in events if e.event_type == "reasoning"]
normal_events = [e for e in events if e.event_type == "normal"]
self.assertEqual(len(reasoning_events), 1)
self.assertGreater(len(normal_events), 0)
def test_empty_input(self):
"""Test handling of empty input."""
parser = HarmonyParser()
events = parser.parse("")
self.assertEqual(len(events), 0)
def test_whitespace_preservation(self):
"""Test that whitespace is preserved correctly."""
parser = HarmonyParser()
text = "<|channel|>analysis<|message|> content with spaces <|end|>"
events = parser.parse(text)
self.assertEqual(len(events), 1)
self.assertEqual(events[0].content, " content with spaces ")
def test_streaming_whitespace_preservation(self):
"""Test that streaming preserves whitespace between chunks."""
parser = HarmonyParser()
# Simulate streaming where space is at chunk boundary
chunks = ["analysis The user typed ", '"wapppa". Not a question.']
all_events = []
for chunk in chunks:
events = parser.parse(chunk)
all_events.extend(events)
# Combine all reasoning content
reasoning_content = "".join(
e.content for e in all_events if e.event_type == "reasoning"
)
# Should preserve the space before the quote
self.assertIn('typed "wapppa"', reasoning_content)
self.assertNotIn(
'typed"wapppa"', reasoning_content
) # Should not be mashed together
def test_consecutive_blocks_same_type(self):
"""Test consecutive blocks of the same type."""
parser = HarmonyParser()
text = (
"<|channel|>analysis<|message|>first reasoning<|end|>"
"<|channel|>analysis<|message|>second reasoning<|end|>"
)
events = parser.parse(text)
self.assertEqual(len(events), 2)
self.assertEqual(events[0].event_type, "reasoning")
self.assertEqual(events[1].event_type, "reasoning")
self.assertEqual(events[0].content, "first reasoning")
self.assertEqual(events[1].content, "second reasoning")
if __name__ == "__main__":
unittest.main()
-603
View File
@@ -1,603 +0,0 @@
import unittest
from sglang.srt.parser.reasoning_parser import (
BaseReasoningFormatDetector,
DeepSeekR1Detector,
KimiDetector,
Qwen3Detector,
ReasoningParser,
StreamingParseResult,
)
from sglang.test.test_utils import CustomTestCase
class TestStreamingParseResult(CustomTestCase):
def test_init_default(self):
"""Test default initialization of StreamingParseResult."""
result = StreamingParseResult()
self.assertEqual(result.normal_text, "")
self.assertEqual(result.reasoning_text, "")
def test_init_with_values(self):
"""Test initialization with specific values."""
result = StreamingParseResult("normal", "reasoning")
self.assertEqual(result.normal_text, "normal")
self.assertEqual(result.reasoning_text, "reasoning")
class TestBaseReasoningFormatDetector(CustomTestCase):
def setUp(self):
self.detector = BaseReasoningFormatDetector(
think_start_token="<think>",
think_end_token="</think>",
force_reasoning=False,
stream_reasoning=True,
)
def test_init(self):
"""Test initialization of BaseReasoningFormatDetector."""
self.assertEqual(self.detector.think_start_token, "<think>")
self.assertEqual(self.detector.think_end_token, "</think>")
self.assertFalse(self.detector._in_reasoning)
self.assertTrue(self.detector.stream_reasoning)
self.assertEqual(self.detector._buffer, "")
self.assertFalse(self.detector.stripped_think_start)
def test_detect_and_parse_normal_text(self):
"""Test parsing normal text without reasoning."""
text = "This is normal text"
result = self.detector.detect_and_parse(text)
self.assertEqual(result.normal_text, text)
self.assertEqual(result.reasoning_text, "")
def test_detect_and_parse_with_start_token(self):
"""Test parsing text starting with think token."""
text = "<think>This is reasoning"
result = self.detector.detect_and_parse(text)
self.assertEqual(result.reasoning_text, "This is reasoning")
self.assertEqual(result.normal_text, "")
def test_detect_and_parse_complete_reasoning(self):
"""Test parsing complete reasoning block."""
text = "<think>This is reasoning</think>This is normal"
result = self.detector.detect_and_parse(text)
self.assertEqual(result.reasoning_text, "This is reasoning")
self.assertEqual(result.normal_text, "This is normal")
def test_detect_and_parse_force_reasoning(self):
"""Test forced reasoning mode."""
detector = BaseReasoningFormatDetector(
"<think>", "</think>", force_reasoning=True
)
text = "This should be reasoning"
result = detector.detect_and_parse(text)
self.assertEqual(result.reasoning_text, "This should be reasoning")
self.assertEqual(result.normal_text, "")
def test_parse_streaming_increment_normal(self):
"""Test streaming parse of normal text."""
result = self.detector.parse_streaming_increment("Hello world")
self.assertEqual(result.normal_text, "Hello world")
self.assertEqual(result.reasoning_text, "")
def test_parse_streaming_increment_partial_token(self):
"""Test streaming parse with partial token."""
# Test partial start token
result = self.detector.parse_streaming_increment("<thi")
self.assertEqual(result.normal_text, "")
self.assertEqual(result.reasoning_text, "")
# Reset detector and test partial end token when in reasoning mode
detector = BaseReasoningFormatDetector("<think>", "</think>")
detector._in_reasoning = True
result = detector.parse_streaming_increment("</thi")
self.assertEqual(result.normal_text, "")
self.assertEqual(result.reasoning_text, "")
def test_parse_streaming_increment_complete_start(self):
"""Test streaming parse with complete start token."""
result = self.detector.parse_streaming_increment("<think>")
self.assertEqual(result.normal_text, "")
self.assertEqual(result.reasoning_text, "")
self.assertTrue(self.detector._in_reasoning)
self.assertTrue(self.detector.stripped_think_start)
def test_parse_streaming_increment_reasoning_content(self):
"""Test streaming parse of reasoning content."""
# First add start token
self.detector.parse_streaming_increment("<think>")
# Then add reasoning content
result = self.detector.parse_streaming_increment("reasoning content")
self.assertEqual(result.reasoning_text, "reasoning content")
self.assertEqual(result.normal_text, "")
def test_parse_streaming_increment_end_token(self):
"""Test streaming parse with end token."""
# Start reasoning mode
self.detector.parse_streaming_increment("<think>")
self.detector.parse_streaming_increment("reasoning")
# End reasoning - the reasoning content accumulated in previous calls is cleared when end token is found
result = self.detector.parse_streaming_increment("</think>normal text")
self.assertEqual(result.reasoning_text, "") # Buffer cleared, returns empty
self.assertEqual(result.normal_text, "normal text")
self.assertFalse(self.detector._in_reasoning)
def test_parse_streaming_increment_no_stream_reasoning(self):
"""Test streaming parse without streaming reasoning."""
detector = BaseReasoningFormatDetector(
"<think>", "</think>", stream_reasoning=False
)
# Start reasoning mode
detector.parse_streaming_increment("<think>")
# Add reasoning content - should not return content
result = detector.parse_streaming_increment("reasoning content")
self.assertEqual(result.reasoning_text, "")
self.assertEqual(result.normal_text, "")
def test_parse_streaming_increment_mixed_content(self):
"""Test streaming parse with mixed content in one chunk."""
result = self.detector.parse_streaming_increment(
"<think>reasoning</think>normal"
)
self.assertEqual(result.reasoning_text, "reasoning")
self.assertEqual(result.normal_text, "normal")
class TestDeepSeekR1Detector(CustomTestCase):
def setUp(self):
self.detector = DeepSeekR1Detector()
def test_init(self):
"""Test DeepSeekR1Detector initialization."""
self.assertEqual(self.detector.think_start_token, "<think>")
self.assertEqual(self.detector.think_end_token, "</think>")
self.assertTrue(self.detector._in_reasoning) # force_reasoning=True
self.assertTrue(self.detector.stream_reasoning)
def test_init_no_stream_reasoning(self):
"""Test DeepSeekR1Detector with stream_reasoning=False."""
detector = DeepSeekR1Detector(stream_reasoning=False)
self.assertFalse(detector.stream_reasoning)
def test_detect_and_parse_r1_format(self):
"""Test parsing DeepSeek-R1 format."""
text = "I need to think about this. The answer is 42."
result = self.detector.detect_and_parse(text)
# Should be treated as reasoning because force_reasoning=True
self.assertEqual(
result.reasoning_text, "I need to think about this. The answer is 42."
)
self.assertEqual(result.normal_text, "")
def test_detect_and_parse_with_end_token(self):
"""Test parsing with end token."""
text = "I think this is the answer</think>The final answer is 42."
result = self.detector.detect_and_parse(text)
self.assertEqual(result.reasoning_text, "I think this is the answer")
self.assertEqual(result.normal_text, "The final answer is 42.")
def test_detect_and_parse_with_start_token(self):
"""Test parsing deepseek-ai/DeepSeek-R1-0528 format, which generates the <think> token."""
text = "<think>I need to think about this.</think>The answer is 42."
result = self.detector.detect_and_parse(text)
# Should be treated as reasoning because force_reasoning=True
self.assertEqual(result.reasoning_text, "I need to think about this.")
self.assertEqual(result.normal_text, "The answer is 42.")
class TestQwen3Detector(CustomTestCase):
def setUp(self):
self.detector = Qwen3Detector()
def test_init(self):
"""Test Qwen3Detector initialization."""
self.assertEqual(self.detector.think_start_token, "<think>")
self.assertEqual(self.detector.think_end_token, "</think>")
self.assertFalse(self.detector._in_reasoning) # force_reasoning=False
self.assertTrue(self.detector.stream_reasoning)
def test_detect_and_parse_qwen3_format(self):
"""Test parsing Qwen3 format."""
text = "<think>Let me think about this problem</think>The answer is 42."
result = self.detector.detect_and_parse(text)
self.assertEqual(result.reasoning_text, "Let me think about this problem")
self.assertEqual(result.normal_text, "The answer is 42.")
def test_detect_and_parse_without_thinking(self):
"""Test parsing without thinking (enable_thinking=False case)."""
text = "Direct answer without thinking."
result = self.detector.detect_and_parse(text)
self.assertEqual(result.normal_text, text)
self.assertEqual(result.reasoning_text, "")
class TestQwen3ForcedReasoningDetector(CustomTestCase):
def setUp(self):
self.detector = Qwen3Detector(force_reasoning=True)
def test_init(self):
"""Test Qwen3ForcedReasoningDetector initialization."""
self.assertEqual(self.detector.think_start_token, "<think>")
self.assertEqual(self.detector.think_end_token, "</think>")
self.assertTrue(self.detector._in_reasoning) # force_reasoning=True
self.assertTrue(self.detector.stream_reasoning)
def test_detect_and_parse_qwen3_forced_reasoning_format(self):
"""Test parsing Qwen3-ForcedReasoning format (no <think> start tag)."""
text = "I need to think about this step by step.</think>The answer is 42."
result = self.detector.detect_and_parse(text)
self.assertEqual(
result.reasoning_text, "I need to think about this step by step."
)
self.assertEqual(result.normal_text, "The answer is 42.")
def test_detect_and_parse_with_start_token(self):
"""Test parsing Qwen3-ForcedReasoning with optional <think> start tag."""
text = "<think>I need to think about this.</think>The answer is 42."
result = self.detector.detect_and_parse(text)
# Should work because base class logic handles both force_reasoning=True OR start token
self.assertEqual(result.reasoning_text, "I need to think about this.")
self.assertEqual(result.normal_text, "The answer is 42.")
def test_streaming_qwen3_forced_reasoning_format(self):
"""Test streaming parse of Qwen3-ForcedReasoning format."""
# First chunk without <think> start
result = self.detector.parse_streaming_increment("I need to")
self.assertEqual(result.reasoning_text, "I need to")
self.assertEqual(result.normal_text, "")
# More reasoning content
result = self.detector.parse_streaming_increment(" think about this.")
self.assertEqual(result.reasoning_text, " think about this.")
self.assertEqual(result.normal_text, "")
# End token with normal text
result = self.detector.parse_streaming_increment("</think>The answer is 42.")
self.assertEqual(result.reasoning_text, "") # Buffer cleared
self.assertEqual(result.normal_text, "The answer is 42.")
class TestKimiDetector(CustomTestCase):
def setUp(self):
self.detector = KimiDetector()
def test_init(self):
"""Test KimiDetector initialization."""
self.assertEqual(self.detector.think_start_token, "◁think▷")
self.assertEqual(self.detector.think_end_token, "◁/think▷")
self.assertFalse(self.detector._in_reasoning)
self.assertTrue(self.detector.stream_reasoning)
def test_detect_and_parse_kimi_format(self):
"""Test parsing Kimi format."""
text = "◁think▷Let me consider this carefully◁/think▷The answer is 42."
result = self.detector.detect_and_parse(text)
self.assertEqual(result.reasoning_text, "Let me consider this carefully")
self.assertEqual(result.normal_text, "The answer is 42.")
def test_detect_and_parse_kimi_no_thinking(self):
"""Test parsing Kimi format without thinking."""
text = "Direct answer without thinking tokens."
result = self.detector.detect_and_parse(text)
self.assertEqual(result.normal_text, text)
self.assertEqual(result.reasoning_text, "")
def test_streaming_kimi_format(self):
"""Test streaming parse of Kimi format."""
# Test partial token
result = self.detector.parse_streaming_increment("◁thi")
self.assertEqual(result.normal_text, "")
self.assertEqual(result.reasoning_text, "")
# Complete start token
result = self.detector.parse_streaming_increment("nk▷Start")
self.assertEqual(result.normal_text, "")
self.assertEqual(result.reasoning_text, "Start")
self.assertTrue(self.detector._in_reasoning)
# Add reasoning content
result = self.detector.parse_streaming_increment("thinking...")
self.assertEqual(result.reasoning_text, "thinking...")
self.assertEqual(result.normal_text, "")
# End token - reasoning content is cleared when end token is processed
result = self.detector.parse_streaming_increment("◁/think▷answer")
self.assertEqual(result.reasoning_text, "") # Buffer cleared
self.assertEqual(result.normal_text, "answer")
class TestReasoningParser(CustomTestCase):
def test_init_valid_model(self):
"""Test initialization with valid model types."""
parser = ReasoningParser("deepseek-r1")
self.assertIsInstance(parser.detector, DeepSeekR1Detector)
parser = ReasoningParser("qwen3")
self.assertIsInstance(parser.detector, Qwen3Detector)
parser = ReasoningParser("kimi")
self.assertIsInstance(parser.detector, KimiDetector)
def test_init_invalid_model(self):
"""Test initialization with invalid model type."""
with self.assertRaises(ValueError) as context:
ReasoningParser("invalid-model")
self.assertIn("Unsupported model type", str(context.exception))
def test_init_no_model(self):
"""Test initialization without model type."""
with self.assertRaises(ValueError) as context:
ReasoningParser(None)
self.assertEqual(str(context.exception), "Model type must be specified")
def test_parse_non_stream(self):
"""Test non-streaming parsing."""
parser = ReasoningParser("qwen3")
reasoning, normal = parser.parse_non_stream(
"<think>Let me think</think>The answer is 42."
)
self.assertEqual(reasoning, "Let me think")
self.assertEqual(normal, "The answer is 42.")
def test_parse_stream_chunk(self):
"""Test streaming chunk parsing."""
parser = ReasoningParser("qwen3")
# First chunk with start token
reasoning, normal = parser.parse_stream_chunk("<think>")
self.assertEqual(reasoning, "")
self.assertEqual(normal, "")
# Second chunk with reasoning content
reasoning, normal = parser.parse_stream_chunk("thinking...")
self.assertEqual(reasoning, "thinking...")
self.assertEqual(normal, "")
# Third chunk with end token and normal text
reasoning, normal = parser.parse_stream_chunk("</think>answer")
self.assertEqual(reasoning, "") # Buffer cleared when end token processed
self.assertEqual(normal, "answer")
def test_case_insensitive_model_type(self):
"""Test case insensitive model type matching."""
parser1 = ReasoningParser("DeepSeek-R1")
parser2 = ReasoningParser("QWEN3")
parser3 = ReasoningParser("Kimi")
self.assertIsInstance(parser1.detector, DeepSeekR1Detector)
self.assertIsInstance(parser2.detector, Qwen3Detector)
self.assertIsInstance(parser3.detector, KimiDetector)
def test_stream_reasoning_parameter(self):
"""Test stream_reasoning parameter is passed correctly."""
parser = ReasoningParser("qwen3", stream_reasoning=False)
self.assertFalse(parser.detector.stream_reasoning)
parser = ReasoningParser("qwen3", stream_reasoning=True)
self.assertTrue(parser.detector.stream_reasoning)
class TestIntegrationScenarios(CustomTestCase):
"""Integration tests for realistic usage scenarios."""
def test_deepseek_r1_complete_response(self):
"""Test complete DeepSeek-R1 response parsing."""
parser = ReasoningParser("deepseek-r1")
text = "I need to solve this step by step. First, I'll analyze the problem. The given equation is x + 2 = 5. To solve for x, I subtract 2 from both sides: x = 5 - 2 = 3.</think>The answer is x = 3."
reasoning, normal = parser.parse_non_stream(text)
self.assertIn("step by step", reasoning)
self.assertIn(
"= 3", reasoning
) # The reasoning contains "x = 5 - 2 = 3" which has "= 3"
self.assertEqual(normal, "The answer is x = 3.")
def test_qwen3_streaming_scenario(self):
"""Test Qwen3 streaming scenario."""
parser = ReasoningParser("qwen3")
chunks = [
"<think>",
"Let me analyze this problem.",
" I need to consider multiple factors.",
"</think>",
"Based on my analysis, the solution is to use a different approach.",
]
all_reasoning = ""
all_normal = ""
for chunk in chunks:
reasoning, normal = parser.parse_stream_chunk(chunk)
all_reasoning += reasoning
all_normal += normal
self.assertIn("analyze", all_reasoning)
self.assertIn("multiple factors", all_reasoning)
self.assertIn("different approach", all_normal)
def test_kimi_streaming_scenario(self):
"""Test Kimi streaming scenario."""
parser = ReasoningParser("kimi")
chunks = [
"◁thi",
"nk▷",
"Let me analyze this problem.",
" I need to consider multiple factors.",
"◁/th",
"ink▷",
"The answer is 42.",
]
all_reasoning = ""
all_normal = ""
for chunk in chunks:
reasoning, normal = parser.parse_stream_chunk(chunk)
all_reasoning += reasoning
all_normal += normal
self.assertIn("analyze", all_reasoning)
self.assertIn("multiple factors", all_reasoning)
self.assertIn("42", all_normal)
def test_empty_reasoning_blocks(self):
"""Test handling of empty reasoning blocks."""
parser = ReasoningParser("qwen3")
text = "<think></think>Just the answer."
reasoning, normal = parser.parse_non_stream(text)
self.assertEqual(reasoning, "")
self.assertEqual(normal, "Just the answer.")
def test_qwen3_forced_reasoning_complete_response(self):
"""Test complete Qwen3-ForcedReasoning response parsing."""
parser = ReasoningParser("qwen3", force_reasoning=True)
text = "Let me solve this step by step. The equation is x + 2 = 5. Subtracting 2 from both sides gives x = 3.</think>The solution is x = 3."
reasoning, normal = parser.parse_non_stream(text)
self.assertIn("step by step", reasoning)
self.assertIn("x = 3", reasoning)
self.assertEqual(normal, "The solution is x = 3.")
def test_qwen3_forced_reasoning_streaming_scenario(self):
"""Test Qwen3-ForcedReasoning streaming scenario."""
parser = ReasoningParser("qwen3", force_reasoning=True)
chunks = [
"I need to analyze",
" this problem carefully.",
" Let me break it down.",
"</think>",
"The final answer is 42.",
]
all_reasoning = ""
all_normal = ""
for chunk in chunks:
reasoning, normal = parser.parse_stream_chunk(chunk)
all_reasoning += reasoning
all_normal += normal
self.assertIn("analyze", all_reasoning)
self.assertIn("break it down", all_reasoning)
self.assertIn("final answer", all_normal)
class TestBufferLossBugFix(CustomTestCase):
"""Test cases for the buffer loss bug fix in parse_streaming_increment."""
def test_partial_end_tag_buffer_loss_bug(self):
"""
Test the bug where partial end tag fragments are lost when followed by normal text.
Bug scenario:
1. _in_reasoning is False
2. new_text is "</" (part of closing thinking tag)
3. Fragment is stored in buffer and empty string is returned
4. Next step: new_text is "answer", _in_reasoning still False
5. Buffer is cleared and "answer" is returned directly
6. The "</" from previous step is lost
This test verifies the fix where line 108 was changed from:
return StreamingParseResult(normal_text=new_text)
to:
return StreamingParseResult(normal_text=current_text)
"""
detector = BaseReasoningFormatDetector("<think>", "</think>")
# Step 1: Send partial end tag when not in reasoning mode
# This should be buffered since it could be start of "</think>"
result1 = detector.parse_streaming_increment("</")
self.assertEqual(result1.normal_text, "")
self.assertEqual(result1.reasoning_text, "")
# Step 2: Send normal text that doesn't complete the end tag
# Before fix: would return only "answer", losing the "</"
# After fix: should return the complete buffered content "</answer"
result2 = detector.parse_streaming_increment("answer")
self.assertEqual(result2.normal_text, "</answer")
self.assertEqual(result2.reasoning_text, "")
def test_partial_start_tag_buffer_preservation(self):
"""
Test that partial start tag fragments are properly preserved.
"""
detector = BaseReasoningFormatDetector("<think>", "</think>")
# Send partial start tag
result1 = detector.parse_streaming_increment("<th")
self.assertEqual(result1.normal_text, "")
self.assertEqual(result1.reasoning_text, "")
# Complete with non-matching text
result2 = detector.parse_streaming_increment("is is text")
self.assertEqual(result2.normal_text, "<this is text")
self.assertEqual(result2.reasoning_text, "")
def test_partial_end_tag_in_reasoning_mode(self):
"""
Test partial end tag handling when already in reasoning mode.
"""
detector = BaseReasoningFormatDetector("<think>", "</think>")
# Enter reasoning mode
detector.parse_streaming_increment("<think>")
detector.parse_streaming_increment("some reasoning")
# Send partial end tag
result1 = detector.parse_streaming_increment("</")
self.assertEqual(result1.normal_text, "")
self.assertEqual(result1.reasoning_text, "")
# Complete the end tag with normal text
result2 = detector.parse_streaming_increment("think>normal text")
self.assertEqual(result2.normal_text, "normal text")
# The reasoning text should be empty since buffer was cleared when end tag was processed
self.assertEqual(result2.reasoning_text, "")
def test_multiple_partial_fragments(self):
"""
Test handling of multiple partial fragments that don't match any tokens.
"""
detector = BaseReasoningFormatDetector("<think>", "</think>")
# Send multiple partial fragments
result1 = detector.parse_streaming_increment("<")
self.assertEqual(result1.normal_text, "")
self.assertEqual(result1.reasoning_text, "")
result2 = detector.parse_streaming_increment("/")
self.assertEqual(result2.normal_text, "")
self.assertEqual(result2.reasoning_text, "")
result3 = detector.parse_streaming_increment("random>")
self.assertEqual(result3.normal_text, "</random>")
self.assertEqual(result3.reasoning_text, "")
def test_edge_case_exact_token_match(self):
"""
Test edge case where buffer content exactly matches a token.
"""
detector = BaseReasoningFormatDetector("<think>", "</think>")
# Build up the exact start token character by character
detector.parse_streaming_increment("<")
detector.parse_streaming_increment("t")
detector.parse_streaming_increment("h")
detector.parse_streaming_increment("i")
detector.parse_streaming_increment("n")
result = detector.parse_streaming_increment("k>")
# Should enter reasoning mode
self.assertEqual(result.normal_text, "")
self.assertEqual(result.reasoning_text, "")
self.assertTrue(detector._in_reasoning)
self.assertTrue(detector.stripped_think_start)
if __name__ == "__main__":
unittest.main()