fix(streaming): keep think tags balanced for interleaved reasoning and exact-match the SSE done marker

This commit is contained in:
mateo-berri 2026-07-17 14:10:49 -04:00
parent 00e0dd1bc1
commit 7ce21e5656
4 changed files with 314 additions and 39 deletions

View file

@ -989,7 +989,7 @@ class CustomStreamWrapper:
)
if _is_delta_empty:
model_response.choices[0].delta = Delta(content=None) # ensure empty delta chunk returned
model_response.choices[0].delta = Delta(content=self._consume_pending_think_close_tag())
# get any function call arguments
model_response.choices[0].finish_reason = map_finish_reason(
finish_reason=self.received_finish_reason
@ -1015,32 +1015,33 @@ class CustomStreamWrapper:
"""
if self.merge_reasoning_content_in_choices is True:
reasoning_content = getattr(model_response.choices[0].delta, "reasoning_content", None)
if reasoning_content:
if self.sent_first_thinking_block is False:
# Ensure content is not None before concatenation
if model_response.choices[0].delta.content is None:
model_response.choices[0].delta.content = ""
model_response.choices[0].delta.content += "<think>" + reasoning_content
self.sent_first_thinking_block = True
elif (
self.sent_first_thinking_block is True
and hasattr(model_response.choices[0].delta, "reasoning_content")
and model_response.choices[0].delta.reasoning_content
):
model_response.choices[0].delta.content = reasoning_content
elif (
self.sent_first_thinking_block is True
and not self.sent_last_thinking_block
and model_response.choices[0].delta.content
):
model_response.choices[0].delta.content = "</think>" + (model_response.choices[0].delta.content or "")
self.sent_last_thinking_block = True
if self.merge_reasoning_content_in_choices is not True:
return
delta = model_response.choices[0].delta
reasoning_content: str = getattr(delta, "reasoning_content", None) or ""
visible_content: str = delta.content or ""
if hasattr(delta, "reasoning_content"):
del delta.reasoning_content
if not reasoning_content and not visible_content:
return
inside_think_block = self.sent_first_thinking_block and not self.sent_last_thinking_block
opening_tag = "<think>" if reasoning_content and not inside_think_block else ""
closing_tag = "</think>" if visible_content and (reasoning_content or inside_think_block) else ""
delta.content = opening_tag + reasoning_content + closing_tag + visible_content
self.sent_first_thinking_block = self.sent_first_thinking_block or bool(reasoning_content)
self.sent_last_thinking_block = self.sent_first_thinking_block and bool(
visible_content or not reasoning_content
)
if hasattr(model_response.choices[0].delta, "reasoning_content"):
del model_response.choices[0].delta.reasoning_content
return
def _consume_pending_think_close_tag(self) -> Optional[str]:
if (
self.merge_reasoning_content_in_choices
and self.sent_first_thinking_block
and not self.sent_last_thinking_block
):
self.sent_last_thinking_block = True
return "</think>"
return None
def _dispatch_provider_chunk(
self,
@ -1678,6 +1679,9 @@ class CustomStreamWrapper:
def finish_reason_handler(self):
model_response = self.model_response_creator()
closing_think_tag = self._consume_pending_think_close_tag()
if closing_think_tag is not None:
model_response.choices[0].delta.content = closing_think_tag
_finish_reason = self.received_finish_reason or self.intermittent_finish_reason
if _finish_reason is not None:
model_response.choices[0].finish_reason = _finish_reason

View file

@ -3,6 +3,7 @@ from abc import abstractmethod
from typing import TYPE_CHECKING, List, Optional, Union, cast
import litellm
from litellm._logging import verbose_logger
if TYPE_CHECKING:
import httpx
@ -63,12 +64,26 @@ def convert_model_response_to_streaming(
raise ValueError(f"Failed to convert ModelResponse to ModelResponseStream: {model_response}. Error: {e}")
MAX_PARTIAL_JSON_LINE_CHARS = 1_000_000
def _try_parse_json_object(payload: str) -> Optional[dict]:
try:
parsed = json.loads(payload)
except json.JSONDecodeError:
return None
if isinstance(parsed, dict):
return parsed
return None
class BaseModelResponseIterator:
def __init__(self, streaming_response, sync_stream: bool, json_mode: Optional[bool] = False):
self.streaming_response = streaming_response
self.response_iterator = self.streaming_response
self.json_mode = json_mode
self.http_response: Optional["httpx.Response"] = None
self.partial_json_line: str = ""
async def aclose(self) -> None:
"""Close the upstream HTTP response so the provider connection is
@ -108,10 +123,37 @@ class BaseModelResponseIterator:
stripped_json_chunk = None
return stripped_json_chunk
def _parse_payload_with_rejoin(self, payload: str) -> Optional[dict]:
rejoined_line = self.partial_json_line + payload
if self.partial_json_line:
rejoined_parsed = _try_parse_json_object(rejoined_line)
if rejoined_parsed is not None:
self.partial_json_line = ""
return rejoined_parsed
parsed = _try_parse_json_object(payload)
if parsed is not None:
if self.partial_json_line:
verbose_logger.debug("Dropping unparseable stream fragment: %s", self.partial_json_line[:1000])
self.partial_json_line = ""
return parsed
if self.partial_json_line:
if len(rejoined_line) > MAX_PARTIAL_JSON_LINE_CHARS:
verbose_logger.debug("Dropping unparseable stream fragment: %s", rejoined_line[:1000])
self.partial_json_line = ""
else:
self.partial_json_line = rejoined_line
return None
if payload.lstrip().startswith("{"):
self.partial_json_line = payload
return None
verbose_logger.debug("Dropping unparseable stream line: %s", payload[:1000])
return None
def _handle_string_chunk(self, str_line: str) -> Union[GenericStreamingChunk, ModelResponseStream]:
# chunk is a str at this point
stripped_json_chunk = BaseModelResponseIterator._string_to_dict_parser(str_line=str_line)
if "[DONE]" in str_line:
payload = litellm.CustomStreamWrapper._strip_sse_data_from_chunk(str_line) or ""
if payload.strip().startswith("[DONE]"):
self.partial_json_line = ""
return GenericStreamingChunk(
text="",
is_finished=True,
@ -120,17 +162,17 @@ class BaseModelResponseIterator:
index=0,
tool_use=None,
)
elif stripped_json_chunk:
return self.chunk_parser(chunk=stripped_json_chunk)
else:
return GenericStreamingChunk(
text="",
is_finished=False,
finish_reason="",
usage=None,
index=0,
tool_use=None,
)
parsed_chunk = self._parse_payload_with_rejoin(payload=payload)
if parsed_chunk:
return self.chunk_parser(chunk=parsed_chunk)
return GenericStreamingChunk(
text="",
is_finished=False,
finish_reason="",
usage=None,
index=0,
tool_use=None,
)
def __next__(self):
while True:

View file

@ -404,6 +404,134 @@ def test_multi_chunk_reasoning_and_content(
assert initialized_custom_stream_wrapper.sent_last_thinking_block is True
def _reasoning_delta_chunk(
content: Optional[str], reasoning_content: Optional[str]
) -> ModelResponseStream:
return ModelResponseStream(
id="chunk-id",
object="chat.completion.chunk",
created=1741037890,
model="glm-5.2",
choices=[
StreamingChoices(
index=0,
delta=Delta(content=content, reasoning_content=reasoning_content),
finish_reason=None,
)
],
)
def _merge_reasoning_delta(
wrapper: CustomStreamWrapper,
content: Optional[str],
reasoning_content: Optional[str],
) -> Optional[str]:
response = _reasoning_delta_chunk(content, reasoning_content)
wrapper._optional_combine_thinking_block_in_choices(response)
assert not hasattr(response.choices[0].delta, "reasoning_content")
return response.choices[0].delta.content
def test_interleaved_reasoning_reopens_think_block(
initialized_custom_stream_wrapper: CustomStreamWrapper,
):
initialized_custom_stream_wrapper.merge_reasoning_content_in_choices = True
deltas_and_expected = [
((None, "plan A"), "<think>plan A"),
((None, " more"), " more"),
(("answer part 1", None), "</think>answer part 1"),
((None, "plan B"), "<think>plan B"),
(("answer part 2", None), "</think>answer part 2"),
]
merged = tuple(
_merge_reasoning_delta(initialized_custom_stream_wrapper, content, reasoning)
for (content, reasoning), _ in deltas_and_expected
)
assert merged == tuple(expected for _, expected in deltas_and_expected)
assert (
"".join(m for m in merged if m)
== "<think>plan A more</think>answer part 1<think>plan B</think>answer part 2"
)
assert initialized_custom_stream_wrapper.sent_last_thinking_block is True
def test_mixed_reasoning_and_content_delta_preserves_content(
initialized_custom_stream_wrapper: CustomStreamWrapper,
):
initialized_custom_stream_wrapper.merge_reasoning_content_in_choices = True
first = _merge_reasoning_delta(
initialized_custom_stream_wrapper, None, "thinking"
)
mixed = _merge_reasoning_delta(
initialized_custom_stream_wrapper, "answer", " final thought"
)
assert first == "<think>thinking"
assert mixed == " final thought</think>answer"
def test_mixed_first_delta_wraps_reasoning_and_keeps_content(
initialized_custom_stream_wrapper: CustomStreamWrapper,
):
initialized_custom_stream_wrapper.merge_reasoning_content_in_choices = True
merged = _merge_reasoning_delta(initialized_custom_stream_wrapper, "hi", "think")
assert merged == "<think>think</think>hi"
assert initialized_custom_stream_wrapper.sent_first_thinking_block is True
assert initialized_custom_stream_wrapper.sent_last_thinking_block is True
def test_stream_end_mid_reasoning_emits_think_close(
initialized_custom_stream_wrapper: CustomStreamWrapper,
):
initialized_custom_stream_wrapper.merge_reasoning_content_in_choices = True
opening = _merge_reasoning_delta(
initialized_custom_stream_wrapper, None, "half-finished thought"
)
assert opening == "<think>half-finished thought"
initialized_custom_stream_wrapper.received_finish_reason = "stop"
final_chunk = ModelResponseStream(
id="chunk-id",
object="chat.completion.chunk",
created=1741037891,
model="glm-5.2",
choices=[
StreamingChoices(index=0, delta=Delta(content=None), finish_reason=None)
],
)
result = initialized_custom_stream_wrapper.return_processed_chunk_logic(
completion_obj={"content": ""},
model_response=final_chunk,
response_obj={},
)
assert result is not None
assert result.choices[0].delta.content == "</think>"
assert result.choices[0].finish_reason == "stop"
assert initialized_custom_stream_wrapper.sent_last_thinking_block is True
def test_finish_reason_handler_closes_open_think_block(
initialized_custom_stream_wrapper: CustomStreamWrapper,
):
initialized_custom_stream_wrapper.merge_reasoning_content_in_choices = True
_merge_reasoning_delta(initialized_custom_stream_wrapper, None, "dangling thought")
final_chunk = initialized_custom_stream_wrapper.finish_reason_handler()
assert final_chunk.choices[0].delta.content == "</think>"
assert final_chunk.choices[0].finish_reason == "stop"
def test_finish_chunk_stays_empty_when_think_block_closed(
initialized_custom_stream_wrapper: CustomStreamWrapper,
):
initialized_custom_stream_wrapper.merge_reasoning_content_in_choices = True
_merge_reasoning_delta(initialized_custom_stream_wrapper, None, "reasoning")
_merge_reasoning_delta(initialized_custom_stream_wrapper, "answer", None)
final_chunk = initialized_custom_stream_wrapper.finish_reason_handler()
assert final_chunk.choices[0].delta.content is None
def test_strip_sse_data_from_chunk():
"""Test the static method that strips 'data: ' prefix from SSE chunks"""
# Test with string inputs

View file

@ -256,3 +256,104 @@ async def test_aclose_is_noop_without_http_response():
)
await iterator.aclose()
class ContentEchoIterator(BaseModelResponseIterator):
def chunk_parser(self, chunk: dict) -> GenericStreamingChunk:
choices = chunk.get("choices") or []
content = choices[0].get("delta", {}).get("content", "") if choices else ""
return GenericStreamingChunk(
text=content or "",
is_finished=False,
finish_reason="",
usage=None,
index=0,
tool_use=None,
)
class TestDoneMarkerExactMatch:
def test_content_containing_done_substring_is_not_treated_as_stream_end(self):
lines = [
'data: {"id":"1","choices":[{"delta":{"content":"foo [DONE] bar"}}]}',
"data: [DONE]",
]
iterator = ContentEchoIterator(streaming_response=iter(lines), sync_stream=True)
chunks = list(iterator)
assert len(chunks) == 2
assert chunks[0]["text"] == "foo [DONE] bar"
assert chunks[0]["is_finished"] is False
assert chunks[1]["is_finished"] is True
assert chunks[1]["finish_reason"] == "stop"
def test_done_line_variants_still_terminate(self):
for line in ["data: [DONE]", "data:[DONE]", "[DONE]", "data: [DONE]"]:
iterator = ContentEchoIterator(
streaming_response=iter([line]), sync_stream=True
)
chunks = list(iterator)
assert len(chunks) == 1
assert chunks[0]["is_finished"] is True
class TestUnicodeLineSeparatorSplitRecovery:
def test_line_split_at_unicode_separator_is_rejoined(self):
lines = [
'data: {"id":"1","choices":[{"delta":{"content":"foo',
'bar"}}]}',
"data: [DONE]",
]
iterator = ContentEchoIterator(streaming_response=iter(lines), sync_stream=True)
chunks = list(iterator)
streamed_text = "".join(chunk["text"] for chunk in chunks)
assert streamed_text == "foobar"
assert all(chunk["is_finished"] is False for chunk in chunks[:-1])
assert chunks[-1]["is_finished"] is True
def test_line_split_into_three_fragments_is_rejoined(self):
lines = [
'data: {"id":"1","choices":[{"delta":{"content":"a',
"b",
'c"}}]}',
"data: [DONE]",
]
iterator = ContentEchoIterator(streaming_response=iter(lines), sync_stream=True)
chunks = list(iterator)
assert "".join(chunk["text"] for chunk in chunks) == "abc"
def test_pending_fragment_does_not_corrupt_next_valid_line(self):
lines = [
'data: {"broken',
'data: {"id":"1","choices":[{"delta":{"content":"ok"}}]}',
"data: [DONE]",
]
iterator = ContentEchoIterator(streaming_response=iter(lines), sync_stream=True)
chunks = list(iterator)
assert "".join(chunk["text"] for chunk in chunks) == "ok"
assert chunks[-1]["is_finished"] is True
@pytest.mark.asyncio
async def test_line_split_at_unicode_separator_is_rejoined_async(self):
async def async_gen():
for line in [
'data: {"id":"1","choices":[{"delta":{"content":"foo',
'bar"}}]}',
"data: [DONE]",
]:
yield line
iterator = ContentEchoIterator(
streaming_response=async_gen(), sync_stream=False
)
chunks = [chunk async for chunk in iterator]
assert "".join(chunk["text"] for chunk in chunks) == "foobar"