diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index e67d0446f23..794e9a4f54c 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -1033,7 +1033,7 @@ class CustomStreamWrapper: visible_content or not reasoning_content ) - def _consume_pending_think_close_tag(self) -> Optional[str]: + def _consume_pending_think_close_tag(self) -> str | None: if ( self.merge_reasoning_content_in_choices and self.sent_first_thinking_block diff --git a/litellm/llms/base_llm/base_model_iterator.py b/litellm/llms/base_llm/base_model_iterator.py index 6f2853851f3..d58124a39c2 100644 --- a/litellm/llms/base_llm/base_model_iterator.py +++ b/litellm/llms/base_llm/base_model_iterator.py @@ -65,9 +65,10 @@ def convert_model_response_to_streaming( MAX_PARTIAL_JSON_LINE_CHARS = 1_000_000 +MAX_PARTIAL_JSON_LINE_FRAGMENTS = 64 -def _try_parse_json_object(payload: str) -> Optional[dict]: +def _try_parse_json_object(payload: str) -> dict | None: try: parsed = json.loads(payload) except json.JSONDecodeError: @@ -84,6 +85,7 @@ class BaseModelResponseIterator: self.json_mode = json_mode self.http_response: Optional["httpx.Response"] = None self.partial_json_line: str = "" + self.partial_json_line_fragments: int = 0 async def aclose(self) -> None: """Close the upstream HTTP response so the provider connection is @@ -123,28 +125,36 @@ 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 + def _drop_partial_json_line(self) -> None: if self.partial_json_line: + verbose_logger.debug("Dropping unparseable stream fragment: %s", self.partial_json_line[:1000]) + self.partial_json_line = "" + self.partial_json_line_fragments = 0 + + def _parse_payload_with_rejoin(self, payload: str) -> dict | None: + rejoined_line = self.partial_json_line + payload + if self.partial_json_line and "}" in payload: rejoined_parsed = _try_parse_json_object(rejoined_line) if rejoined_parsed is not None: self.partial_json_line = "" + self.partial_json_line_fragments = 0 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 = "" + self._drop_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: + if ( + len(rejoined_line) <= MAX_PARTIAL_JSON_LINE_CHARS + and self.partial_json_line_fragments < MAX_PARTIAL_JSON_LINE_FRAGMENTS + ): self.partial_json_line = rejoined_line - return None - if payload.lstrip().startswith("{"): + self.partial_json_line_fragments += 1 + return None + self._drop_partial_json_line() + if payload.lstrip().startswith("{") and len(payload) <= MAX_PARTIAL_JSON_LINE_CHARS: self.partial_json_line = payload + self.partial_json_line_fragments = 1 return None verbose_logger.debug("Dropping unparseable stream line: %s", payload[:1000]) return None @@ -153,7 +163,7 @@ class BaseModelResponseIterator: # chunk is a str at this point payload = litellm.CustomStreamWrapper._strip_sse_data_from_chunk(str_line) or "" if payload.strip().startswith("[DONE]"): - self.partial_json_line = "" + self._drop_partial_json_line() return GenericStreamingChunk( text="", is_finished=True, diff --git a/tests/test_litellm/llms/base_llm/test_base_model_iterator.py b/tests/test_litellm/llms/base_llm/test_base_model_iterator.py index 8f23dbdd2e3..24be0588def 100644 --- a/tests/test_litellm/llms/base_llm/test_base_model_iterator.py +++ b/tests/test_litellm/llms/base_llm/test_base_model_iterator.py @@ -340,6 +340,61 @@ class TestUnicodeLineSeparatorSplitRecovery: assert "".join(chunk["text"] for chunk in chunks) == "ok" assert chunks[-1]["is_finished"] is True + def test_new_json_object_arriving_at_capacity_blowout_is_seeded(self): + from litellm.llms.base_llm.base_model_iterator import MAX_PARTIAL_JSON_LINE_CHARS + + oversized_head = 'data: {"pad":"' + "x" * (MAX_PARTIAL_JSON_LINE_CHARS - 8) + lines = [ + oversized_head, + 'data: {"id":"1","choices":[{"delta":{"content":"after reset', + '"}}]}', + "data: [DONE]", + ] + iterator = ContentEchoIterator(streaming_response=iter(lines), sync_stream=True) + + chunks = list(iterator) + + assert "".join(chunk["text"] for chunk in chunks) == "after reset" + assert chunks[-1]["is_finished"] is True + + def test_fragment_buffer_is_bounded_by_fragment_count(self): + from litellm.llms.base_llm.base_model_iterator import MAX_PARTIAL_JSON_LINE_FRAGMENTS + + lines = ( + ['data: {"choices":[{"delta":{"content":"runaway'] + + ["a"] * MAX_PARTIAL_JSON_LINE_FRAGMENTS + + [ + 'z"}}]}', + 'data: {"id":"1","choices":[{"delta":{"content":"ok"}}]}', + "data: [DONE]", + ] + ) + iterator = ContentEchoIterator(streaming_response=iter(lines), sync_stream=True) + + chunks = list(iterator) + + streamed_text = "".join(chunk["text"] for chunk in chunks) + assert "runaway" not in streamed_text + assert streamed_text == "ok" + + def test_oversized_object_start_is_not_buffered(self): + from litellm.llms.base_llm.base_model_iterator import MAX_PARTIAL_JSON_LINE_CHARS + + oversized_start = 'data: {"choices":[{"delta":{"content":"huge' + "x" * MAX_PARTIAL_JSON_LINE_CHARS + lines = [ + oversized_start, + '"}}]}', + 'data: {"id":"1","choices":[{"delta":{"content":"ok"}}]}', + "data: [DONE]", + ] + iterator = ContentEchoIterator(streaming_response=iter(lines), sync_stream=True) + + chunks = list(iterator) + + streamed_text = "".join(chunk["text"] for chunk in chunks) + assert "huge" not in streamed_text + assert streamed_text == "ok" + @pytest.mark.asyncio async def test_line_split_at_unicode_separator_is_rejoined_async(self): async def async_gen():