mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(streaming): keep think tags balanced for interleaved reasoning and exact-match the SSE done marker
This commit is contained in:
parent
00e0dd1bc1
commit
7ce21e5656
4 changed files with 314 additions and 39 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue