From 4b6d8704720d0ef659040ff63a76e4ca13809522 Mon Sep 17 00:00:00 2001 From: Pradeep Ramola Date: Sun, 27 Sep 2026 19:02:19 -0400 Subject: [PATCH 1/3] fix(streaming): reject incomplete generic chunks --- litellm/litellm_core_utils/streaming_handler.py | 5 +++-- .../test_streaming_overhead.py | 16 ++++------------ 2 files changed, 7 insertions(+), 14 deletions(-) diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index d2853a625c9..c24b67ac7ec 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -54,7 +54,8 @@ FUNCTION_CALL_ATTRIBUTE: Final = "function_call" _SYNC_ITER_EXHAUSTED: Final = object() -_GCHUNK_FIELDS: Final[frozenset] = frozenset(GChunk.__annotations__) +_GCHUNK_FIELDS: Final[frozenset[str]] = frozenset(GChunk.__annotations__) +_GCHUNK_REQUIRED_FIELDS: Final[frozenset[str]] = frozenset(GChunk.__required_keys__) _USAGE_COST_HEADER_PROVIDERS: Final[frozenset[str]] = frozenset({LlmProviders.OPENROUTER.value}) @@ -2467,7 +2468,7 @@ def generic_chunk_has_all_required_fields(chunk: dict) -> bool: :param chunk: The dictionary to check. :return: True if all required fields are present, False otherwise. """ - return all(key in _GCHUNK_FIELDS for key in chunk) + return _GCHUNK_REQUIRED_FIELDS <= chunk.keys() <= _GCHUNK_FIELDS def convert_generic_chunk_to_model_response_stream( diff --git a/tests/unit/litellm_core_utils/test_streaming_overhead.py b/tests/unit/litellm_core_utils/test_streaming_overhead.py index 8fb0659ab5a..f56c57f6c47 100644 --- a/tests/unit/litellm_core_utils/test_streaming_overhead.py +++ b/tests/unit/litellm_core_utils/test_streaming_overhead.py @@ -126,25 +126,17 @@ def test_gchunk_fields_is_frozenset(): assert _GCHUNK_FIELDS == frozenset(GChunk.__annotations__) -def test_generic_chunk_has_all_required_fields_uses_module_constant(monkeypatch): - """generic_chunk_has_all_required_fields must use _GCHUNK_FIELDS, not __annotations__. - - The check semantics: every key in `chunk` must be a known GChunk field. - This identifies GChunk-shaped dicts (all keys are valid GChunk fields). - """ +def test_generic_chunk_has_all_required_fields_rejects_incomplete_chunks(): valid_chunk = _make_generic_chunk("hello") assert generic_chunk_has_all_required_fields(valid_chunk) is True - # A dict with an extra unknown key should return False — the unknown key - # is not a GChunk field, so the chunk is not a pure GChunk. extra_key_chunk = dict(valid_chunk) extra_key_chunk["unknown_extra_key"] = "value" assert generic_chunk_has_all_required_fields(extra_key_chunk) is False - # A dict with only known GChunk fields but fewer keys still passes because - # all its keys are valid (subset of GChunk fields). - partial_chunk = {"text": "hi", "is_finished": False} - assert generic_chunk_has_all_required_fields(partial_chunk) is True + for required_field in GChunk.__required_keys__: + incomplete_chunk = {key: value for key, value in valid_chunk.items() if key != required_field} + assert generic_chunk_has_all_required_fields(incomplete_chunk) is False # --------------------------------------------------------------------------- From b594c88fe5be65ae533a014c742afa467ac2f118 Mon Sep 17 00:00:00 2001 From: Pradeep Ramola Date: Wed, 30 Sep 2026 00:40:43 -0400 Subject: [PATCH 2/3] fix(streaming): handle incomplete chunks safely --- .../litellm_core_utils/streaming_handler.py | 22 ++++++---- .../test_streaming_handler.py | 40 +++++++++++++++++++ .../test_streaming_overhead.py | 3 ++ 3 files changed, 58 insertions(+), 7 deletions(-) diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index c24b67ac7ec..52cf323febf 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -1110,11 +1110,10 @@ class CustomStreamWrapper: completion_obj: dict[str, Any], ) -> _ProviderChunkResult: response_obj: dict[str, Any] = {} - if ( - isinstance(chunk, ModelResponseStream) - and self.custom_llm_provider is not None - and self.custom_llm_provider in litellm._custom_providers - ): + is_registered_custom_provider: Final = ( + self.custom_llm_provider is not None and self.custom_llm_provider in litellm._custom_providers + ) + if isinstance(chunk, ModelResponseStream) and is_registered_custom_provider: _has_content: Final = bool( chunk.choices and chunk.choices[0].delta is not None @@ -1133,10 +1132,19 @@ class CustomStreamWrapper: chunk.choices[0].finish_reason = None return _ProviderChunkEarlyReturn(chunk) + is_generic_chunk: Final = isinstance(chunk, dict) and generic_chunk_has_all_required_fields(chunk=chunk) if ( isinstance(chunk, dict) - and generic_chunk_has_all_required_fields(chunk=chunk) # check if chunk is a generic streaming chunk - ) or (self.custom_llm_provider and self.custom_llm_provider in litellm._custom_providers): + and not is_generic_chunk + and (chunk.keys() <= _GCHUNK_FIELDS or is_registered_custom_provider) + ): + partial_chunk_usage: Final = chunk.get("usage") + if isinstance(partial_chunk_usage, dict): + model_response.usage = litellm.Usage(**partial_chunk_usage) + return _ProviderChunkParsed(cast(dict[str, object], chunk)) + return _ProviderChunkEarlyReturn(None) + + if is_generic_chunk: if self.received_finish_reason is not None: _chunk_has_content: Final = isinstance(chunk, dict) and ( bool(chunk.get("text", "")) diff --git a/tests/unit/litellm_core_utils/test_streaming_handler.py b/tests/unit/litellm_core_utils/test_streaming_handler.py index d07e8822eb0..7b5bcaa262b 100644 --- a/tests/unit/litellm_core_utils/test_streaming_handler.py +++ b/tests/unit/litellm_core_utils/test_streaming_handler.py @@ -2511,6 +2511,46 @@ def test_usage_only_chunk_not_dropped_when_finish_reason_already_set( assert result.usage is not None +@pytest.mark.parametrize("custom_llm_provider", [None, "my-custom-llm"]) +@pytest.mark.parametrize( + "chunk", + [ + pytest.param({}, id="empty"), + pytest.param({"is_finished": False}, id="finish-state-only"), + ], +) +def test_chunk_creator_skips_incomplete_generic_chunks( + monkeypatch: pytest.MonkeyPatch, + initialized_custom_stream_wrapper: CustomStreamWrapper, + custom_llm_provider: str | None, + chunk: dict[str, object], +): + if custom_llm_provider is not None: + monkeypatch.setattr(litellm, "_custom_providers", [custom_llm_provider]) + initialized_custom_stream_wrapper.custom_llm_provider = custom_llm_provider + + result = initialized_custom_stream_wrapper.chunk_creator(chunk=chunk) + + assert result is None + assert initialized_custom_stream_wrapper.chunks == [] + + +@pytest.mark.parametrize("custom_llm_provider", [None, "my-custom-llm"]) +def test_chunk_creator_records_incomplete_usage_chunk( + monkeypatch: pytest.MonkeyPatch, + initialized_custom_stream_wrapper: CustomStreamWrapper, + custom_llm_provider: str | None, +): + if custom_llm_provider is not None: + monkeypatch.setattr(litellm, "_custom_providers", [custom_llm_provider]) + initialized_custom_stream_wrapper.custom_llm_provider = custom_llm_provider + + result = initialized_custom_stream_wrapper.chunk_creator(chunk={"usage": {"prompt_tokens": 1}}) + + assert result is None + assert initialized_custom_stream_wrapper.chunks[-1].usage.prompt_tokens == 1 + + def _run_dispatch(wrapper: CustomStreamWrapper, chunk): model_response = wrapper.model_response_creator() completion_obj = {"content": ""} diff --git a/tests/unit/litellm_core_utils/test_streaming_overhead.py b/tests/unit/litellm_core_utils/test_streaming_overhead.py index f56c57f6c47..5bab5d7a7a5 100644 --- a/tests/unit/litellm_core_utils/test_streaming_overhead.py +++ b/tests/unit/litellm_core_utils/test_streaming_overhead.py @@ -138,6 +138,9 @@ def test_generic_chunk_has_all_required_fields_rejects_incomplete_chunks(): incomplete_chunk = {key: value for key, value in valid_chunk.items() if key != required_field} assert generic_chunk_has_all_required_fields(incomplete_chunk) is False + for degenerate_chunk in ({}, {"is_finished": False}, {"usage": {"prompt_tokens": 1}}): + assert generic_chunk_has_all_required_fields(degenerate_chunk) is False + # --------------------------------------------------------------------------- # 2. Cached model name and provider at init time From fce092a885cc0b0159600d821a6aa45baf71bc0a Mon Sep 17 00:00:00 2001 From: Pradeep Ramola Date: Wed, 30 Sep 2026 18:28:32 -0400 Subject: [PATCH 3/3] fix(streaming): preserve complete custom chunks --- .../litellm_core_utils/streaming_handler.py | 5 +++- .../test_streaming_handler.py | 28 +++++++++++++++++++ 2 files changed, 32 insertions(+), 1 deletion(-) diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index 52cf323febf..4ea7327bdfc 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -1132,7 +1132,10 @@ class CustomStreamWrapper: chunk.choices[0].finish_reason = None return _ProviderChunkEarlyReturn(chunk) - is_generic_chunk: Final = isinstance(chunk, dict) and generic_chunk_has_all_required_fields(chunk=chunk) + is_generic_chunk: Final = isinstance(chunk, dict) and ( + generic_chunk_has_all_required_fields(chunk=chunk) + or (is_registered_custom_provider and _GCHUNK_REQUIRED_FIELDS <= chunk.keys()) + ) if ( isinstance(chunk, dict) and not is_generic_chunk diff --git a/tests/unit/litellm_core_utils/test_streaming_handler.py b/tests/unit/litellm_core_utils/test_streaming_handler.py index 7b5bcaa262b..cc68551fdca 100644 --- a/tests/unit/litellm_core_utils/test_streaming_handler.py +++ b/tests/unit/litellm_core_utils/test_streaming_handler.py @@ -2551,6 +2551,34 @@ def test_chunk_creator_records_incomplete_usage_chunk( assert initialized_custom_stream_wrapper.chunks[-1].usage.prompt_tokens == 1 +def test_custom_provider_complete_generic_chunk_with_extra_fields_is_preserved( + monkeypatch: pytest.MonkeyPatch, +): + custom_llm_provider = "my-custom-llm" + monkeypatch.setattr(litellm, "_custom_providers", [custom_llm_provider]) + wrapper = CustomStreamWrapper( + completion_stream=iter( + [ + { + "text": "hello", + "is_finished": True, + "finish_reason": "stop", + "usage": None, + "custom_metadata": {"trace_id": "trace-1"}, + } + ] + ), + model="custom-model", + logging_obj=MagicMock(), + custom_llm_provider=custom_llm_provider, + ) + + chunks = list(wrapper) + + assert "".join(chunk.choices[0].delta.content or "" for chunk in chunks) == "hello" + assert chunks[-1].choices[0].finish_reason == "stop" + + def _run_dispatch(wrapper: CustomStreamWrapper, chunk): model_response = wrapper.model_response_creator() completion_obj = {"content": ""}